From 0219f2791ea31adea366dd3a739967eb9c1ce6e2 Mon Sep 17 00:00:00 2001 From: Gaspard Zoss Date: Fri, 28 Aug 2026 05:22:12 -0700 Subject: [PATCH] Add custom Mitsuba integrator for GNM half-Lambert shading. This change introduces GnmHalfLambertIntegrator, a custom Mitsuba integrator that implements a half-Lambert shading model combined with a Cook-Torrance specular lobe. This allows Mitsuba renders to closely match the output of the PyRender-based render_gnm pipeline. The integrator is dynamically registered with Mitsuba via a variant callback. PiperOrigin-RevId: 972526336 --- gnm/shape/data/versions/v3_0/gnm_head.npz | Bin 53305389 -> 53305389 bytes .../visualization/integrators/__init__.py | 56 ++ .../gnm_half_lambert_integrator.py | 189 +++++ gnm/shape/visualization/render_common.py | 681 ++++++++++++++++++ gnm/shape/visualization/render_common_test.py | 351 +++++++++ gnm/shape/visualization/render_gnm.py | 650 +---------------- gnm/shape/visualization/render_gnm_test.py | 188 +---- 7 files changed, 1310 insertions(+), 805 deletions(-) create mode 100644 gnm/shape/visualization/integrators/__init__.py create mode 100644 gnm/shape/visualization/integrators/gnm_half_lambert_integrator.py create mode 100644 gnm/shape/visualization/render_common.py create mode 100644 gnm/shape/visualization/render_common_test.py diff --git a/gnm/shape/data/versions/v3_0/gnm_head.npz b/gnm/shape/data/versions/v3_0/gnm_head.npz index 4ef4591a1546e9478facdc2317911841d6d8ce84..5511c703b69a57a7762def3a6701960b0abeba37 100644 GIT binary patch delta 5612 zcmajj1yEFt0)}DOV_>5qii)Brim-}`VvF6a*d2)7?QeHrx31mYfns-ecXxOE-^+64 zUXNa#d7s&t^X=IM_Ux>?V&?|86)JAUGDb8P2PY>dha%<2#`SkC5H_|&YzGeq{y91% zbExRnyhOZ!enEqR95)16K7FzjDL=W=XG``!3~jTkS&a&lAAPn2u`Kk(;^$iMms2`l z`L0vi)->=k5)Ya7#i9%qO+{BRR7~ZhVyW0Fj*6?|srV{^a#jgdBE@-JR1)Q?lB#6N zO(j<;R7&NpQmNF+L#0t^m8VLl(km~OL1k3lDwFb2zRIdHt1QY-WmVZ!cIB^fsGKT5 z)l3De<|;(BP%Tv})mpVtZB;wfUUg6% zRVUS1bx~baH`QJBP(4*I)m#0e`l!CDpX#p$sDWyb8mxw>p=y{Ku12VlYLptSLe&^G zR*h5R)dV$BO;VH96g5>%Q`6N9HB-$}HZ@z#QFGNiHD4`I3)LdESS?XY)iSkQtxzk~ zDz#dzQESyYwO(yd8`UPYS#42U)ixETwyPZ~TYW5l#Cj-d)HU_dm84ly7m zI6*9k4RIhY#Dn;d0GuHqB!a}?0!hFXl0q_YgXE9`Qi40Ag4EyvX&^0lLOMtfUXTGY zf;VIWAMgb$WQHu@2U#H-WCwr90XZQ6azSp$19>4IAhX&9P8bM=d0!^VA1VeKO zffmpbT0v`Q18t!lw1*DR5jsI<=mK4#8+3;r&=Yz=Z}M99wxv$4SQfO?1TMq01m<- zI1ESNC>(?1Z~{)kDL4&h;4GYj^Kbz!!X>y2SKumKgX?euZo)0N4R_!!+=Kh@03O04 zcnnYADLjMc@B&`KD|iiW;4Qp^_wWHe!YB9)U*IczgYWReQEz`1-Eag87!VDjLkx%s zP7n)XLmY?;@gP1V0B1-Di6Ak!Si<%8XNf*lAH{5QCb(L>O`pX)EI;&FEc~Vp<}#Qi zpWJ zGl`^-4BQ|&q=1y*4yhnDct9FR3!ab;(t{UdfQ;Y`nZO5p!3von3;01+$OhTLA96rW z2!LFW8}dM2$OrkM02G8kCOftn2lb%=G=xUb7@9y+Xa>R1973Q4w1igB8rncxXb0_~19XH=&>6Zw zSLg=ap$GJYUeFug|G+~!xC5u%V0UIfR(TcR>K-t3+rG#Y=Dih z2{ywP*b3Vq47S4#2#1}p3wFaE*bDn$KOBIAa0m{=5jYCR;5eLslW+=7!x=aW=ioeC zfQxVmF2fbL3fJH|+<=>K3vR<5xC{5-K0JVj@CY8m6L<>G;5od2m+%T+!y9-D@8CUr zfRFGAKEoII3g6&6{IHCux82YUN3eha(I7g+fSBL}u^=|Yfw&M4;zI&(hJ=s^5`zmQ z0ar*0$-oVgLkdV~gzN3k2#UrV>Be2|hL`D$bmtu2NF)B|oza$jx4Y4m?RecVzmfmf z*4yT|$K9yFwz6)R-=J^4;Jo4Eb@22&X2HR!3=h-8`>kur`}j;944)aC+L*Y8t-qUx z;c42ge$tGm8u4HM(7u|8G?3N^@9XAiw3y3%pKkLs(wXk-d_I`R@e1=aL^>nfD=fWH zc0PCLq^}Y(nC@`%R0wzA!9)hg2;Pthe2nll4SbC|KAdFINMFNhI?2nAoFx2`4hC8* zMa$b#g!&p8T?2plPUB>XR?KSZVClTf79939yiK#0ArZ5?I(V-svm<_1Bkdn%k$+L2 ztolZUZ7dr#z`iWwz;ZTJ<9N^@sVN?O06~u+0@< zHFB7?feAYBgAdTb-R{gv=nJ4Cf0#u)d_R8Ko)R%&H_Ucau==oe>xUiRxa>Ji6imD%)&_RBbG^;WB_rd8wx{PTI$ zu5C54nO62oHR@b@w8~*x*{{#2RpvTYBd=*?zW^gwdVSfn%4vFi*{{E-Rh+u~wSRqK z+b_SURX45t{<^xZRW8#i@*=VSCiJ|@*8APLeqF8AD*s>4mA?M}xrTl|;w-Hyna*Xu PAR~T87n@rHtKsr5^kZLo delta 5612 zcmajj1yEFNABJ(*V_+Aes8}eXum&n(cNZ46Vt03A2P$^!+TGpQ9oXI7-SyrV%Lngx z^!=QX^VxGBK$<%{K<&(tbkEjj-(Xq#KZYE+*38g9S20vf<)mV%*eZ^StKzBnDuGI<5~;+Bb(~dF<)V_Q z#szR!;Dx!+2Vyd_*p-QS!sf+gTB%m4 z)oP7etJbOYYJ=LSHmS{Oi`uHTsqJcq+NpM_-6~Y=QG3-swO<`j2h|~USRGME)iHHk zolqy$DRo+fsWa-VI;YO73+ke}q%Nx~>Z-b?uB#jBrn;qWt2^qhx~J}|2kN1Eq#mm$ z>Zy9Bo~sw?rFx}at2gScdZ*s259*`(q&}-J>Z|&uzN;VVr}`B<*vtfDhz^T#y@lArJULUdRXep#T(wLQoirKv5_L#i0b0gi=r% z%0O8t2j!sxRD?=U8T_FNRE26#9cn-T)P!148|pw^s0a0-0W^d_XatR+2{eUf&>UJo zOK1hHp$)W!cF-O=Ku72VouLbKg>KLtdO%O;1-+pU^o4%V9|k}W41_^27>2-57zV>( z1dN1H5DcSX42*>k7zg8F0!)NSFd3%6RG0?SVFt{ESuh*sfDPutJeUs)U?D7m#jpgH z!ZKJ6D_|w8g4M7F*1|ei4;x@3Y=X_O1-8OA*bX~jC+vdV5DI%>FYJT;Z~zX%Avg?2 z;3yn}<8T5_!YMcnVQ>b{!Z|n(7vLgXg3E9PuEI6A4maQ?+=AP12kyc>xDOBDAv}V| z@C2U1Gk6X!;3d3**YF13!aH~mAK)W=g3s^;zQQ;74nN>0{BqRg&!Pv8U;zW7L3D@# zF~JF9L2QTvaUmYWhXjxi5t|v5uAbE*5%Q2mK7j zQcQEPWc%A#LmhmUnlZOsE*3Xit1T{;_@?2Y?%`{W(ZR9N|2OPq8Y}X(M8(&>Aq}*W zNCwHl6;eP-NCj??8qz>oNC)Y`9Wp>h@PJI<37H`ac!4)qAuD8q?2rR|ASdL4+~5m& zzz_05KFALRpdb{2!cYW?LNO=~C7>jfg3?e1%0f9P4;7#yRD#Oj4^^NlRD*d{_VrVG%5bC9o8h!E#suD`6F^hBdGj*1>w%02^TwY=$kc z6}G{4*a16X7wm>m*aLfEAMA$%a1ai`VK@Ru;TRl;6L1nv!D$GCGjJBp!Fjj<7vU0I zhAVItuEBM<0XN|m+=e@F7w*A*cmNOK5j=(`@D!fGb9ezS;T61wH}DqT!F%`sAK?>x zhA;3HzQK3+0YBlFCAh&ZLk}Fm0tQ5b=nw;9f)m7o*boQeLOh5M2_PXPg2a#noFOT= zKr%=Uu8;y!LMkIvmp>yQ8W+;dd)*8V(}mPOT}fR?!~gR3w58bRX7pfoq8^wpjXjsvejx&>cIlZs5b4JS9hcJe4hJ!hr5x%^j!Z5r8$qsPIp6OFhV_cW;7}+ z|@B~gKGJz*#hAiM^gsuzpHtu<`$du9EhSjvl?KdnEdPN6|SQtwg z;%#`E#@44gRk`uKP0b5{_FJPHL{3VzdTkW zx9Ri_6a9HdgLQD98)Kz)1E~04#==iOP`8k5?g3V#*x&9UIo7cnshOIosiOZqr86e3 zX*CL&&KPl*$-b7RhG{B~Y0Ccn#}+yDPE!T{?^N?zRwJ)z%Kp74YOM=@oU(uIiJB@= z+iK)8trf9PBVX5QO_^WU?^EgYA7_4D_I(_6#&Md;YC5BR8%IsO*HjMERKy03d|h?x zS`8o5lzpp4t+iiM<{usV{*0Q+TF+|aH?3vgfKgL6P5GLp?E5cjDo*`BecJZz7d6#W zQ|9}V5xXz)J>JxmpXrPd8!B?DLW4i8_50L%O_@K(hz%FHR>p?^v42=vR PWcYh@wz&pc4d;IVl_sP< diff --git a/gnm/shape/visualization/integrators/__init__.py b/gnm/shape/visualization/integrators/__init__.py new file mode 100644 index 00000000..5f2a188b --- /dev/null +++ b/gnm/shape/visualization/integrators/__init__.py @@ -0,0 +1,56 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Module that collects and registers all custom GNM integrators.""" + +from __future__ import annotations + +import functools +import importlib + +import mitsuba as mi # pyrefly: ignore[missing-import] + + +def _integrators_variant_callback(old: str | None, new: str) -> None: + """Imports and registers all custom integrators. + + This follows Mitsuba's structure to allow changing the variant for custom + plugins. + + Args: + old: The old variant. + new: The new variant. + """ + del old # unused + + if new is None or new.startswith('scalar'): + return + + # pylint: disable=g-import-not-at-top,import-outside-toplevel + from gnm.shape.visualization.integrators import gnm_half_lambert_integrator + # pylint: enable=g-import-not-at-top,import-outside-toplevel + + importlib.reload(gnm_half_lambert_integrator) + gnm_half_lambert_integrator.register() + + +@functools.cache +def register() -> None: + """Registers all custom GNM integrators. + + If integrators were already registered, this is a no-op. + """ + mi.detail.add_variant_callback(_integrators_variant_callback) + if variant := mi.variant(): + _integrators_variant_callback(None, variant) diff --git a/gnm/shape/visualization/integrators/gnm_half_lambert_integrator.py b/gnm/shape/visualization/integrators/gnm_half_lambert_integrator.py new file mode 100644 index 00000000..0a7954cb --- /dev/null +++ b/gnm/shape/visualization/integrators/gnm_half_lambert_integrator.py @@ -0,0 +1,189 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Custom Mitsuba integrator implementing GNM half-Lambert shading. + +This integrator aims to reproduce the shading of render_gnm (which uses +PyRender) so that Mitsuba renders match that pipeline as closely as possible. + +The class subclasses `mi.SamplingIntegrator` at module top level, which requires +a Mitsuba variant to already be set. This module is therefore only imported +lazily (after `mi.set_variant(...)`) by the package `__init__.py`'s variant +callback. Callers should invoke the package-level `register()` (see +`__init__.py`) before loading a scene that uses this integrator. +""" + +from __future__ import annotations + +import drjit as dr # pyrefly: ignore[missing-import] +import mitsuba as mi # pyrefly: ignore[missing-import] +import numpy as np + +# Light intensity empirically chosen to match render_gnm. +_LIGHT_INTENSITY = 3.34 + + +class GnmHalfLambertIntegrator(mi.SamplingIntegrator): + """Integrator that implements GNM half-Lambert shading and flat shading.""" + + def __init__(self, props: mi.Properties): + """Initializes the integrator. + + Args: + props: Mitsuba properties. + """ + super().__init__(props) + self.light_dir = mi.ScalarVector3f(props.get('light_dir', [1.0, 1.0, 1.0])) + self.light_intensity = mi.Float( + props.get( + 'light_intensity', + _LIGHT_INTENSITY, + ) # pyrefly: ignore[bad-argument-type] + ) + self.include_shading = bool(props.get('include_shading', True)) + + def sample( + self, + scene: mi.Scene, + sampler: mi.Sampler, + ray: mi.Ray3f, + medium: mi.Medium | None = None, + active: mi.Bool | bool = True, + ) -> tuple[mi.Spectrum, mi.Bool, list[mi.Float]]: + """Samples the integrator. + + Args: + scene: The Mitsuba scene. + sampler: The Mitsuba sampler. + ray: The Mitsuba ray. + medium: The Mitsuba medium. + active: Whether the ray is active. + + Returns: + A tuple of the radiance, whether the ray is valid, and the list of + intermediate depths. + """ + del sampler, medium + scene_intersection = scene.ray_intersect(ray, mi.Bool(active)) + is_valid = active & scene_intersection.is_valid() + + has_vertex_colors = scene_intersection.shape.has_attribute('vertex_colors') + vertex_colors = dr.select( + has_vertex_colors, + scene_intersection.shape.eval_attribute( + 'vertex_colors', scene_intersection, is_valid & has_vertex_colors + ), + mi.Color3f(1.0, 1.0, 1.0), + ) + + diffuse_reflectance = scene_intersection.bsdf().eval_diffuse_reflectance( + scene_intersection, is_valid + ) + base_color = vertex_colors * diffuse_reflectance + + if self.include_shading: + radiance = vertex_colors * self._shade_half_lambert( + scene_intersection=scene_intersection, + diffuse_reflectance=diffuse_reflectance, + ) + else: + radiance = dr.power( # pyrefly: ignore[unsupported-operation] + base_color, 2.2 + ) + + radiance = dr.select(is_valid, radiance, mi.Color3f(0.0, 0.0, 0.0)) + return mi.Spectrum(radiance), is_valid, [] + + def _shade_half_lambert( + self, + scene_intersection: mi.SurfaceInteraction3f, + diffuse_reflectance: mi.Color3f | mi.Spectrum, + ) -> mi.Color3f: + """Computes reflected radiance using the GNM half-Lambert shading model. + + Combines a half-Lambert diffuse term with a GLTF metallic-roughness + specular lobe (Cook-Torrance) so that Mitsuba renders match the shading of + the render_gnm pipeline. + + The shading math is evaluated in the local Mitsuba shading frame defined by + `scene_intersection.sh_frame`, following idiomatic Mitsuba BSDF/integrator + conventions. Both the light and view directions are transformed into that + orthonormal frame, where the shading normal is (0, 0, 1). Each `N . x` term + therefore reduces to the z-component of the transformed vector, exposed via + `mi.Frame3f.cos_theta`. Because `sh_frame.to_local` is an orthonormal + transform it preserves lengths and dot products, so this formulation is + numerically equivalent to evaluating the same dot products in world space + (any difference is at floating-point epsilon from the extra transform). + + Args: + scene_intersection: Surface interaction at the shading point. Its + `sh_frame` defines the local shading frame and `wi` is the view + direction (`-ray.d`) already expressed in that frame. + diffuse_reflectance: Diffuse reflectance (base color) at the intersection. + + Returns: + The reflected radiance, before modulation by vertex colors. + """ + # Transform the (world-space) light direction into the local shading frame; + # the view direction is already available there as `wi` (== -ray.d). + light_direction_local = scene_intersection.to_local( + mi.Vector3f(self.light_dir) + ) + view_direction_local = scene_intersection.wi + half_vector_local = dr.normalize( + light_direction_local + view_direction_local + ) + + # In the local frame the normal is (0, 0, 1), so N.L and N.V are just the + # z-components (cos_theta). Half-Lambert-style wrap of N.L (via the 0.8/1.8 + # constants) softens the terminator; n_dot_v and v_dot_h are the usual + # clamped cosine terms. + cos_theta_light = mi.Frame3f.cos_theta(light_direction_local) + cos_theta_view = mi.Frame3f.cos_theta(view_direction_local) + n_dot_l = dr.clip((cos_theta_light + 0.8) / 1.8, 0.0, 1.0) + n_dot_v = dr.clip(dr.abs(cos_theta_view), 0.001, 1.0) + v_dot_h = dr.clip(dr.dot(view_direction_local, half_vector_local), 0.0, 1.0) + + # GLTF PBR specular (Cook-Torrance) with roughness=1.0, metallic=0.0: + # Schlick Fresnel, Smith geometry term, and a constant (roughness=1) NDF. + fresnel = 0.04 + mi.Float(0.96) * ( + dr.power(1.0 - v_dot_h, 5.0) # pyrefly: ignore[unsupported-operation] + ) + geometry_view = n_dot_v / (0.5 * n_dot_v + 0.5) + geometry_light = n_dot_l / (0.5 * n_dot_l + 0.5) + geometry = geometry_view * geometry_light + distribution = 1.0 / np.pi + + diffuse_color = mi.Color3f(diffuse_reflectance * 0.96) + diffuse_contrib = mi.Color3f( + (1.0 - fresnel) + * diffuse_color + / np.pi # pyrefly: ignore[unsupported-operation] + ) + + spec_contrib = mi.Color3f( + (fresnel * geometry * distribution) + / ( + 4.0 * n_dot_l * n_dot_v + 0.001 + ) # pyrefly: ignore[unsupported-operation] + ) + + return mi.Color3f( + n_dot_l * self.light_intensity * (diffuse_contrib + spec_contrib) + ) + + +def register() -> None: + """Registers the GNM half-Lambert integrator plugin.""" + mi.register_integrator('gnm_half_lambert', GnmHalfLambertIntegrator) diff --git a/gnm/shape/visualization/render_common.py b/gnm/shape/visualization/render_common.py new file mode 100644 index 00000000..0eb165c1 --- /dev/null +++ b/gnm/shape/visualization/render_common.py @@ -0,0 +1,681 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Shared helpers for render_gnm backends (Pyrender, etc.). + +This module factors out the backend-agnostic pieces that are common to the +different `render_gnm` implementations: default scene parameters, texture +loading, camera/projection helpers and batching utilities. +""" + +from collections.abc import Sequence +import functools + +from etils import epath +from gnm.shape import gnm_numpy +from gnm.shape.visualization import camera_conversions +import imageio +import immutabledict +import numpy as np +import numpy.typing as npt + +FloatArray = npt.NDArray[np.floating] +ColorOrImage = npt.NDArray[np.uint8] | FloatArray | Sequence[float] | float + +_pkg = __package__ or 'gnm.shape.visualization' +_TEXTURES_DIR = epath.resource_path(_pkg).parent / 'data' / 'textures' +EDGEFLOW_TEXTURE_BY_BODY_PART = immutabledict.immutabledict({ + gnm_numpy.GNMBodyPart.HEAD: str(_TEXTURES_DIR / 'edgeflow_bw_4k.png'), +}) + +# Default parameters for scene. +DEFAULT_IMAGE_SIZE = (240, 320) +DEFAULT_CAMERA_DISTANCE = 2.0 +DEFAULT_NEAR = 0.01 +DEFAULT_FAR = 100.0 +DEFAULT_TARGET_FILL_FACTOR = 0.4 +DEFAULT_BACKGROUND_COLOR = (0.95, 0.95, 0.95) + + +# Placeholder for lazy loading the default texture. +class DefaultTexture: + """Placeholder for lazy loading the default texture.""" + + +DEFAULT_TEXTURE = DefaultTexture() +Texture = FloatArray | DefaultTexture | dict[str, FloatArray] | None + + +def project_points_for_gnm( + gnm_np: gnm_numpy.GNM, + points_world: np.ndarray | None = None, + vertices: np.ndarray | None = None, + world_to_camera: FloatArray | None = None, + camera_to_image: FloatArray | None = None, + image_size: tuple[int, int] = DEFAULT_IMAGE_SIZE, + multiple_gnms: bool = False, + **kwargs, +) -> np.ndarray: + """Projects world points under the same conditions as a render_gnm call. + + Intended for identifying the per-frame positions of 3D points (e.g. GNM + joints) in the same reference space used by render_gnm. + + For a description of the other arguments, please see the docstring of + `render_gnm`. Any shading related arguments are ignored. + + Args: + gnm_np: The GNM model. + points_world: The world-space points to project, (..., P, 3). Defaults to + the vertex positions of the template GNM. + vertices: The GNM vertices in world space, (..., V, 3). If not provided, + will use the template vertices. Used for default camera setup. + world_to_camera: The world-to-camera transformation, (..., 4, 4). + camera_to_image: The camera-to-image transformation, (..., 4, 4). + image_size: The width and height of the rendered image in pixels: (W, H). + multiple_gnms: If True, vertices is expected to be shape (..., M, V, 3), and + we render M GNMs per image. Default cameras will be set relative to the + first GNM in the sequence. + **kwargs: Any additional arguments expected for render_gnm (ignored). + + Returns: + The projected points in image space, (..., P, 2). + """ + + del kwargs + + if points_world is None: + points_world = gnm_np.template_vertex_positions + points_world = points_world.astype(np.float32) + + if vertices is None: + vertices = gnm_np.template_vertex_positions + + if not multiple_gnms: + vertices = vertices[..., None, :, :] # Inject 'M' dimension. + elif (vertices_dim := vertices.ndim) < 3: + raise ValueError( + f'Called with {multiple_gnms=}, but vertices is only {vertices_dim}D.' + ) + + # Define default camera params based on the first GNM in the 'M' dimension. + vertices_for_cameras = vertices[..., 0, :, :] + + if world_to_camera is None: + world_to_camera = get_look_at_world_to_camera( + gnm_np, + vertices_for_cameras, + ) + + if camera_to_image is None: + camera_to_image = get_fill_factor_camera_to_image( + gnm_np, vertices_for_cameras, image_size=image_size + ) + + # Convert from OpenCV to OpenGL convention. + world_to_camera = camera_conversions.opencv_extrinsics_to_opengl( + world_to_camera + ) + camera_to_image = ( + camera_conversions.opencv_intrinsics_matrix_to_opengl_view_matrix( + camera_to_image, + width=image_size[0], + height=image_size[1], + near=DEFAULT_NEAR, + far=DEFAULT_FAR, + ) + ) + + # Find the maximum batch dimension that satisfies all batch-able arguments. + try: + batch_dims = get_batch_dim( + (points_world, 2), + (vertices, 3), + (world_to_camera, 2), + (camera_to_image, 2), + ) + except ValueError as e: + raise ValueError( + f' Batch dimensions incompatible: points_world {points_world.shape},' + f' vertices {vertices.shape}, world_to_camera {world_to_camera.shape},' + f' camera_to_image {camera_to_image.shape}.' + ) from e + + def batchify(arr, non_batch_dims): + """Broadcast to batch dimensions, and flatten the batch dimensions.""" + arr = np.broadcast_to(arr, (*batch_dims, *arr.shape[-non_batch_dims:])) + return arr.reshape(int(np.prod(batch_dims)), *arr.shape[-non_batch_dims:]) + + points_world = batchify(points_world, 2) + world_to_camera = batchify(world_to_camera, 2) + camera_to_image = batchify(camera_to_image, 2) + + # Perform projection. + view_projection_matrix = camera_to_image @ world_to_camera + + homogenous_ones = np.ones( + (*points_world.shape[:-1], 1), dtype=points_world.dtype + ) + points_homogeneous = np.concatenate([points_world, homogenous_ones], axis=-1) + points_clip_space = ( + view_projection_matrix[:, None, :, :] @ points_homogeneous[..., None] + ) + points_clip_space = points_clip_space[..., 0] + + # Perform perspective division to get Normalized Device Coordinates (NDC). + points_ndc = points_clip_space[..., :3] / points_clip_space[..., [3]] + + # Convert NDC to image space. + # NDC range is [-1, 1]. Image space range is [0, image_size]. + width, height = image_size + points_image_space = ( + (points_ndc[..., :2] + 1.0) * 0.5 * np.array([width, height]) + ) + + # Flip +Y up renders to +Y down. + points_image_space[..., 1] = height - points_image_space[..., 1] + + return points_image_space.reshape(*batch_dims, *points_image_space.shape[-2:]) + + +def get_look_at_world_to_camera( + gnm_np: gnm_numpy.GNM, + vertices_world: np.ndarray, + azimuthal_angle: np.ndarray | float = 0.0, + polar_angle: np.ndarray | float = 0.0, + camera_distance: np.ndarray | float | None = DEFAULT_CAMERA_DISTANCE, + share_camera: np.ndarray | bool = True, + y_up: np.ndarray | bool = False, + look_at_vertex_groups: Sequence[str] = ('hockey_mask',), + left_vertex_groups: Sequence[str] = ('ears', '&left'), + right_vertex_groups: Sequence[str] = ('ears', '&right'), + forward_vertex_groups: Sequence[str] = ('nose_region',), +) -> np.ndarray: + """Compute world-to-camera matrices for a 'look-at' transform. + + Returns matrices in OpenCV convention. + + Args: + gnm_np: The GNM model. + vertices_world: The GNM vertices in world space, (..., V, 3). + azimuthal_angle: The azimuthal angle of the camera in degrees, (..., 1). + polar_angle: The polar angle of the camera in degrees, (..., 1). + camera_distance: The distance of the camera from the head, (..., 1). + share_camera: Whether to use the first frame's vertices only for camera + generation, (..., 1). It is assumed that the first dimension of vertices + is the time dimension. + y_up: Whether to use the Y-up convention for the world space, (..., 1). + look_at_vertex_groups: The vertex groups to look at. + left_vertex_groups: The vertex groups to use for the left axis. + right_vertex_groups: The vertex groups to use for the right axis. + forward_vertex_groups: The vertex groups to use for the forward axis. + + Returns: + The world-to-camera matrices, (..., 4, 4). + """ + + batch_dims = vertices_world.shape[:-2] + + if not batch_dims: + # If there is no batch dimension, we don't need to share the camera. + share_camera = False + + if camera_distance is None: + camera_distance = DEFAULT_CAMERA_DISTANCE + + azimuthal_angle = _adjust_scalar_shape(azimuthal_angle, batch_dims) + polar_angle = _adjust_scalar_shape(polar_angle, batch_dims) + camera_distance = _adjust_scalar_shape(camera_distance, batch_dims) + share_camera = _adjust_scalar_shape(share_camera, batch_dims) + y_up = _adjust_scalar_shape(y_up, batch_dims) + + first_frame_vertices = np.broadcast_to( + vertices_world[:1], vertices_world.shape + ) + + vertices_for_camera = np.where( + share_camera[..., None], first_frame_vertices, vertices_world + ) + + gnm_axes = _get_gnm_axes( + vertices_for_camera, + left_vertex_groups=left_vertex_groups, + right_vertex_groups=right_vertex_groups, + forward_vertex_groups=forward_vertex_groups, + gnm_np=gnm_np, + ) + camera_target = _vertex_group_mean( + vertices_for_camera, look_at_vertex_groups, gnm_np + ) + camera_location = camera_target + _get_camera_offset( + gnm_axes, + camera_distance, + azimuthal_angle, + polar_angle, + ) + + y_up_vector = np.array([0.0, 1.0, 0.0]) + up_vector = np.where(y_up, y_up_vector, gnm_axes[1]) + world_to_camera = right_handed_look_at( + camera_location, camera_target, up_vector + ) + + world_to_camera = camera_conversions.opengl_extrinsics_to_opencv( + world_to_camera + ) + + return world_to_camera + + +def get_fill_factor_camera_to_image( + gnm_np: gnm_numpy.GNM, + vertices: np.ndarray, + target_fill_factor: np.ndarray | float = DEFAULT_TARGET_FILL_FACTOR, + camera_distance: np.ndarray | float = DEFAULT_CAMERA_DISTANCE, + image_size: tuple[int, int] = DEFAULT_IMAGE_SIZE, + vertex_group_name: str = 'hockey_mask', +) -> np.ndarray: + """Compute the camera-to-image matrix for a 'fill factor' transform. + + The camera intrinsics are determined to fill the width of the given vertex + group with the image. Returns matrices in OpenCV convention. + + Args: + gnm_np: The GNM model. + vertices: The GNM vertices in world space, (..., V, 3). + target_fill_factor: The desired fill factor of the GNM mesh in the image, + (..., 1). + camera_distance: The distance of the camera from the head, (..., 1). + image_size: The width and height of the rendered image in pixels: (W, H). + vertex_group_name: The vertex group to use for the projection. + + Returns: + The camera-to-image matrix, (..., 4, 4). + """ + batch_dims = vertices.shape[:-2] + target_fill_factor = _adjust_scalar_shape(target_fill_factor, batch_dims) + camera_distance = _adjust_scalar_shape(camera_distance, batch_dims) + + width, height = image_size + + # Get the maximum width of the vertex group. + vertex_group_indices = gnm_np.vertex_group_indices(vertex_group_name) + vertex_group_points = vertices[..., vertex_group_indices, :] + mask_width = np.ptp(vertex_group_points, axis=-2)[..., :1] + dimension = mask_width / target_fill_factor + + vertical_field_of_view = np.atan2(dimension / 2.0, camera_distance) * 2 + aspect_ratio = _adjust_scalar_shape(width / height, batch_dims) + + near = _adjust_scalar_shape(DEFAULT_NEAR, batch_dims) + far = _adjust_scalar_shape(DEFAULT_FAR, batch_dims) + + camera_to_image = _right_handed_perspective( + vertical_field_of_view=vertical_field_of_view, + aspect_ratio=aspect_ratio, + near=near, + far=far, + ) + + camera_to_image = camera_conversions.opengl_intrinsics_to_opencv_matrix( + camera_to_image, width=width, height=height + ) + + return camera_to_image + + +def get_spin_world_to_camera( + gnm_np: gnm_numpy.GNM, + vertices: np.ndarray, + has_time_dimension: bool = False, + spin_period: int = 60, + spin_azimuth_limit: float = 20.0, + spin_polar_limit: float = 5.0, + num_frames: int | None = None, + **kwargs, +) -> np.ndarray: + """Compute world-to-camera matrices for a 'spin' transform. + + Args: + gnm_np: The GNM model. + vertices: The GNM vertices in world space, (..., V, 3). + has_time_dimension: Whether the first dimension of the vertices is time. If + false, will broadcast vertices over the spin period. + spin_period: The number of frames in a full spin. + spin_azimuth_limit: The maximum azimuthal angle of the camera in degrees. + spin_polar_limit: The maximum polar angle of the camera in degrees. + num_frames: Optional number of frames to generate (defaults to spin_period). + **kwargs: Additional arguments to pass to get_look_at_world_to_camera. + + Returns: + The world-to-camera matrices, (spin_period, ..., 4, 4) if has_time_dimension + is False, otherwise (..., 4, 4). + """ + + if num_frames is not None: + if not has_time_dimension and ( + vertices.ndim < 3 or vertices.shape[0] != num_frames + ): + vertices = np.broadcast_to(vertices, (num_frames, *vertices.shape)) + elif has_time_dimension: + num_frames = vertices.shape[0] + else: + num_frames = spin_period + vertices = np.broadcast_to(vertices, (num_frames, *vertices.shape)) + + batch_dims = vertices.shape[:-2] + + # Spin frequency is determined by spin_period. + max_angle = np.pi * 2 * (num_frames / spin_period) + angles = np.linspace(-0, max_angle, num_frames) + + azimuthal_angle = np.sin(angles)[..., None] * spin_azimuth_limit + polar_angle = np.cos(angles)[..., None] * spin_polar_limit + + # Broadcast azimuthal and polar angles to the batch dimensions. + azimuthal_angle = np.broadcast_to(azimuthal_angle, (*batch_dims, 1)) + polar_angle = np.broadcast_to(polar_angle, (*batch_dims, 1)) + + return get_look_at_world_to_camera( + gnm_np, + vertices, + azimuthal_angle=azimuthal_angle, + polar_angle=polar_angle, + **kwargs, + ) + + +def _get_gnm_axes( + vertices: np.ndarray, + left_vertex_groups: Sequence[str], + right_vertex_groups: Sequence[str], + forward_vertex_groups: Sequence[str], + gnm_np: gnm_numpy.GNM, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Determines the right, up, and forwards direction vectors of GNM. + + Args: + vertices: The GNM vertices in world space, (..., V, 3). + left_vertex_groups: The vertex groups to use for the left axis. + right_vertex_groups: The vertex groups to use for the right axis. + forward_vertex_groups: The vertex groups to use for the forward axis. + gnm_np: The GNM model. + + Returns: + A tuple of right, up, and forwards direction vectors, each (..., 3). + """ + + left = _vertex_group_mean(vertices, left_vertex_groups, gnm_np) + right = _vertex_group_mean(vertices, right_vertex_groups, gnm_np) + forward_point = _vertex_group_mean(vertices, forward_vertex_groups, gnm_np) + back_point = (left + right) / 2.0 + right = right - left + right = right / np.linalg.norm(right, axis=-1, keepdims=True) + forwards = forward_point - back_point + forwards = forwards / np.linalg.norm(forwards, axis=-1, keepdims=True) + up = np.linalg.cross(right, forwards) + return right, up, forwards + + +def right_handed_look_at( + camera_position: np.ndarray, look_at: np.ndarray, up_vector: np.ndarray +) -> np.ndarray: + """Builds a right handed look at view matrix. + + Args: + camera_position: The position of the camera, (..., 3). + look_at: The position of the target, (..., 3). + up_vector: The up vector of the camera, (..., 3). + + Returns: + The right handed look at view matrix, (..., 4, 4). + """ + + z_axis = look_at - camera_position + z_axis /= np.linalg.norm(z_axis, axis=-1, keepdims=True) + + horizontal_axis = np.cross(z_axis, up_vector, axis=-1) + horizontal_axis /= np.linalg.norm(horizontal_axis, axis=-1, keepdims=True) + + vertical_axis = np.cross(horizontal_axis, z_axis, axis=-1) + + def _dot_last_axis(arr1, arr2): + return np.einsum('...i,...i->...', arr1, arr2)[..., None] + + batch_shape = horizontal_axis.shape[:-1] + zeros = np.zeros((*batch_shape, 3), dtype=horizontal_axis.dtype) + ones = np.ones((*batch_shape, 1), dtype=horizontal_axis.dtype) + + tx = -_dot_last_axis(horizontal_axis, camera_position) + ty = -_dot_last_axis(vertical_axis, camera_position) + tz = _dot_last_axis(z_axis, camera_position) + + row1 = np.concatenate([horizontal_axis, tx], axis=-1) + row2 = np.concatenate([vertical_axis, ty], axis=-1) + row3 = np.concatenate([-z_axis, tz], axis=-1) + row4 = np.concatenate([zeros, ones], axis=-1) + matrix = np.stack([row1, row2, row3, row4], axis=-2) + return matrix + + +def _get_camera_offset( + head_axes: tuple[np.ndarray, np.ndarray, np.ndarray], + camera_distance: np.ndarray, + azimuthal_angle: np.ndarray, + polar_angle: np.ndarray, +) -> np.ndarray: + """Determines how to place the camera so it points at the face. + + Args: + head_axes: Three-tuple representing normalised right, up, and forwards face + direction vectors, each (..., 3). + camera_distance: Camera distance from the face (meters), (..., 1). + azimuthal_angle: Places camera to the left or right, in degrees, (..., 1). + polar_angle: Places camera above or below the head, in degrees, (..., 1). + + Returns: + A batch of offset vectors to apply to camera targets (..., 3). + """ + right, up, forwards = head_axes + theta = np.deg2rad(90 + azimuthal_angle) + phi = np.deg2rad(90 + polar_angle) + x = camera_distance * np.sin(phi) * np.cos(theta) * right + y = camera_distance * np.sin(phi) * np.sin(theta) * forwards + z = camera_distance * np.cos(phi) * up + return x + y + z + + +def _right_handed_perspective( + vertical_field_of_view: np.ndarray, + aspect_ratio: np.ndarray, + near: np.ndarray, + far: np.ndarray, +) -> np.ndarray: + """Builds a right handed perspective projection matrix. + + Similar to tensorflow_graphics.rendering.camera.perspective.right_handed. + + Args: + vertical_field_of_view: The vertical field of view in radians, (..., 1). + aspect_ratio: The aspect ratio of the image, (..., 1). + near: The near clipping plane distance, (..., 1). + far: The far clipping plane distance, (..., 1). + + Returns: + The perspective projection matrix, (..., 4, 4) in OpenGL convention. + """ + + itan_half_vertical_field_of_view = 1.0 / np.tan(vertical_field_of_view * 0.5) + zero = np.zeros_like(itan_half_vertical_field_of_view) + one = np.ones_like(itan_half_vertical_field_of_view) + near_minus_far = near - far + + row1 = np.concatenate( + (itan_half_vertical_field_of_view / aspect_ratio, zero, zero, zero), + axis=-1, + ) + row2 = np.concatenate( + (zero, itan_half_vertical_field_of_view, zero, zero), axis=-1 + ) + row3 = np.concatenate( + ( + zero, + zero, + (far + near) / near_minus_far, + 2.0 * far * near / near_minus_far, + ), + axis=-1, + ) + row4 = np.concatenate((zero, zero, -one, zero), axis=-1) + + matrix = np.stack((row1, row2, row3, row4), axis=-2) + return matrix + + +@functools.cache +def load_edgeflow_texture(texture_path: str | None) -> FloatArray | None: + """Loads the edgeflow texture from a given path. + + Args: + texture_path: The path to the edgeflow texture. + + Returns: + The edgeflow texture as a float32 array, or None if texture_path is None. + """ + if texture_path is None: + return None + with epath.Path(texture_path).open('rb') as f: + image = imageio.imread(f).astype(np.float32) + image = (image / np.iinfo(np.uint8).max) * 0.5 + 0.5 + image.flags.writeable = False + return image + + +def _vertex_group_mean( + vertices: np.ndarray, + group_names: Sequence[str], + gnm_np: gnm_numpy.GNM, +): + """Gets the average point of GNM vertex groups. + + Args: + vertices: The GNM vertices in world space, (..., V, 3). + group_names: The names of the vertex groups. + gnm_np: The GNM model. + + Returns: + The average point of the vertex groups, (..., 3). + """ + indices = gnm_np.vertex_group_indices(*group_names) + return vertices[..., indices, :].mean(axis=-2) + + +def load_texture( + gnm_np: gnm_numpy.GNM, + texture: Texture = DEFAULT_TEXTURE, +) -> dict[str, npt.NDArray[np.uint8]]: + """Loads the texture as a (potentially batched) image. + + Args: + gnm_np: The GNM model. + texture: The texture to load. If DEFAULT_TEXTURE, will load the edgeflow + texture. If an ndarray, will use the given texture for skin. If a dict, + will use the given texture for each part. If None, will use a white + (plain) texture. + + Returns: + The texture image, [0-255] uint8, (..., H, W, 3). + """ + texture_dict = {} + if texture is DEFAULT_TEXTURE: + edgeflow_path = EDGEFLOW_TEXTURE_BY_BODY_PART.get(gnm_np.body_part) + if edgeflow_path is not None: + edgeflow = load_edgeflow_texture(edgeflow_path) + if edgeflow is not None: + texture_dict['skin'] = edgeflow[..., None] + elif isinstance(texture, np.ndarray): + texture_dict['skin'] = texture + elif isinstance(texture, dict): + texture_dict = texture + + # Fill remaining parts with white texture. + for component in gnm_np.mesh_component_names: + if component not in texture_dict: + texture_dict[component] = np.ones((64, 64, 1)).astype(np.float32) + + def _to_3channel_uint8(texture_image: np.ndarray) -> npt.NDArray[np.uint8]: + if texture_image.dtype != np.uint8: + texture_image = (texture_image * 255.0).astype(np.uint8) + if texture_image.shape[-1] == 1: + texture_image = np.repeat(texture_image, 3, axis=-1) + return texture_image + + texture_dict = { + part: _to_3channel_uint8(texture_dict[part]) for part in texture_dict + } + + return texture_dict + + +def _adjust_scalar_shape( + array: np.ndarray | float, batch_dims: Sequence[int] +) -> np.ndarray: + """Adjusts the shape of a scalar to match the batch dimensions. + + Args: + array: The array to adjust. + batch_dims: The batch dimensions to match. + + Returns: + The adjusted array, (..., 1). + """ + + # If the array is a scalar or 1D array, add a final dimension of size 1. + array = np.array(array) + if array.ndim <= 1: + array = array[..., None] + + return np.broadcast_to(array, (*batch_dims, 1)) + + +def get_batch_dim(*arrays: tuple[np.ndarray | None, int]) -> tuple[int, ...]: + """Finds the largest batch dimension that all arrays can be broadcast to. + + Requires that all arrays have batch dimensions that can be broadcast together. + + e.g. If given an array [B, C, 4] and an array [A, B, C, 3], return (A, B, C). + + Args: + *arrays: The arrays to expand. Tuples of (array, non_batch_dims), where + non_batch_dims is the number of rightmost dimensions of the array that are + not batch dimensions. If None, will be ignored. + + Returns: + The batch dimensions that all arrays can be safely broadcast to. + Raises: + ValueError: If the arrays cannot be broadcast together. + """ + batch_dims: list[tuple[int, ...]] = [] + for array, non_batch_dims in arrays: + if array is None: + continue + batch_dims.append(array.shape[:-non_batch_dims]) + + max_batch_dims = max(batch_dims, key=len) + + # Verify that all arrays can be broadcast together. + for batch_dim in batch_dims: + np.broadcast_shapes(batch_dim, max_batch_dims) + + return max_batch_dims diff --git a/gnm/shape/visualization/render_common_test.py b/gnm/shape/visualization/render_common_test.py new file mode 100644 index 00000000..25e16df2 --- /dev/null +++ b/gnm/shape/visualization/render_common_test.py @@ -0,0 +1,351 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for backend-agnostic helpers in render_common.""" + +from absl.testing import absltest +from absl.testing import parameterized +from gnm.shape import gnm_numpy +from gnm.shape.data.versions import gnm_test_catalog +from gnm.shape.visualization import camera_conversions +from gnm.shape.visualization import render_common +import numpy as np + +_TupleOfInts = tuple[int, ...] + + +def _get_random_parameters( + batch_dims: _TupleOfInts, + gnm_np: gnm_numpy.GNM, +) -> dict[str, np.ndarray]: + """Returns random GNM parameters. + + Args: + batch_dims: The batch dimensions. + gnm_np: The GNM model. + + Returns: + A dictionary of random GNM parameters. + """ + identity = np.random.uniform( + -1.5, 1.5, size=batch_dims + (gnm_np.identity_dim,) + ) + expression = np.random.uniform( + -1.5, 1.5, size=batch_dims + (gnm_np.expression_dim,) + ) + rotations = np.random.uniform( + -0.2, 0.2, size=batch_dims + (gnm_np.num_joints, 3) + ) + translation = np.random.uniform(-0.5, 0.5, size=batch_dims + (3,)) * 0.0 + return dict( + identity=identity.astype(np.float32), + expression=expression.astype(np.float32), + rotations=rotations.astype(np.float32), + translation=translation.astype(np.float32), + ) + + +class TestGetBatchDim(parameterized.TestCase): + """Tests for get_batch_dim.""" + + def test_get_batch_dim(self): + """Tests that get_batch_dim returns the correct batch dimensions.""" + a, b, c = 1, 2, 3 + array_a = np.zeros((a, b, c, 10)) + array_b = None + array_c = np.zeros((b, c, 5, 5)) + batch_dims = render_common.get_batch_dim( + (array_a, 1), (array_b, 2), (array_c, 2) + ) + self.assertEqual(batch_dims, (a, b, c)) + + def test_get_batch_dim_raises_error(self): + """Tests that get_batch_dim raises an error for incompatible dimensions.""" + array_a = np.zeros((1, 2, 3, 10)) + array_b = np.zeros((4, 5, 6, 5)) + with self.assertRaises(ValueError): + render_common.get_batch_dim((array_a, 1), (array_b, 1)) + + +class TestGetLookAtWorldToCamera(parameterized.TestCase): + """Tests for get_look_at_world_to_camera.""" + + gnms: dict[str, gnm_numpy.GNM] + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.gnms = {} + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS: + cls.gnms[version] = gnm_numpy.GNM.from_local( + gnm_numpy.GNMMajorVersion(version.removeprefix('v')), + gnm_numpy.GNMVariant.HEAD, + ) + + @parameterized.named_parameters(*[ + (version, version) + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS + ]) + def test_basic(self, version): + """Tests that get_look_at_world_to_camera returns the correct matrix.""" + gnm_np = self.gnms[version] + vertices = gnm_np.template_vertex_positions + world_to_camera_opencv = render_common.get_look_at_world_to_camera( + gnm_np=gnm_np, + vertices_world=vertices, + ) + world_to_camera_opengl = camera_conversions.opencv_extrinsics_to_opengl( + world_to_camera_opencv + ) + self.assertEqual(world_to_camera_opengl.shape, (4, 4)) + + with self.subTest('Approximately identity rotation.'): + np.testing.assert_allclose( + world_to_camera_opengl[:3, :3], np.eye(3), atol=0.02 + ) + + with self.subTest('Translated from hockey mask.'): + hockey_mask_indices = gnm_np.vertex_group_indices('hockey_mask') + hockey_mask_z = vertices[hockey_mask_indices, 2].mean() + self.assertAlmostEqual( + -world_to_camera_opengl[2, 3], + hockey_mask_z + render_common.DEFAULT_CAMERA_DISTANCE, + delta=0.01, + ) + + @parameterized.named_parameters(*[ + (version, version) + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS + ]) + def test_share_camera_no_batch(self, version): + """Tests that share_camera=False has no effect for no batch dimensions.""" + gnm_np = self.gnms[version] + parameters = _get_random_parameters((), gnm_np) + vertices = gnm_np(**parameters) + world_to_camera_no_share = render_common.get_look_at_world_to_camera( + gnm_np=gnm_np, + vertices_world=vertices, + share_camera=False, + ) + + world_to_camera_share = render_common.get_look_at_world_to_camera( + gnm_np=gnm_np, + vertices_world=vertices, + share_camera=True, + ) + + np.testing.assert_allclose(world_to_camera_no_share, world_to_camera_share) + + @parameterized.named_parameters(*[ + (version, version) + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS + ]) + def test_share_camera(self, version): + """Tests that share_camera=True/False have the correct effect.""" + gnm_np = self.gnms[version] + parameters = _get_random_parameters((10,), gnm_np) + vertices = gnm_np(**parameters) + + world_to_camera_0 = render_common.get_look_at_world_to_camera( + gnm_np=gnm_np, vertices_world=vertices[0] + ) + + world_to_camera_share = render_common.get_look_at_world_to_camera( + gnm_np=gnm_np, + vertices_world=vertices, + share_camera=True, + ) + + world_to_camera_no_share = render_common.get_look_at_world_to_camera( + gnm_np=gnm_np, + vertices_world=vertices, + share_camera=False, + ) + + with self.subTest('Shared camera matches first camera.'): + np.testing.assert_allclose( + world_to_camera_share, + np.broadcast_to(world_to_camera_0, world_to_camera_share.shape), + ) + + with self.subTest('Unshared cameras are all different.'): + self.assertFalse( + np.allclose( + world_to_camera_no_share, + np.broadcast_to( + world_to_camera_0, world_to_camera_no_share.shape + ), + ) + ) + + +class TestGetFillFactorCameraToImage(parameterized.TestCase): + """Tests for get_fill_factor_camera_to_image.""" + + gnms: dict[str, gnm_numpy.GNM] + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.gnms = {} + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS: + cls.gnms[version] = gnm_numpy.GNM.from_local( + gnm_numpy.GNMMajorVersion(version.removeprefix('v')), + gnm_numpy.GNMVariant.HEAD, + ) + + @parameterized.named_parameters(*[ + (version, version) + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS + ]) + def test_basic(self, version): + """Tests that get_fill_factor_camera_to_image returns the correct matrix.""" + gnm_np = self.gnms[version] + vertices = gnm_np.template_vertex_positions + camera_to_image_opencv = render_common.get_fill_factor_camera_to_image( + gnm_np=gnm_np, + vertices=vertices, + ) + camera_to_image_opengl = ( + camera_conversions.opencv_intrinsics_matrix_to_opengl_view_matrix( + camera_to_image_opencv, + width=320, + height=240, + near=0.1, + far=100.0, + ) + ) + self.assertEqual(camera_to_image_opengl.shape, (4, 4)) + with self.subTest('Lower triangular is all zeros.'): + submatrix = camera_to_image_opengl[:3, :3] + self.assertTrue((np.tril(submatrix, k=-1) == 0.0).all()) + + +class TestLoadTexture(parameterized.TestCase): + """Tests for load_texture.""" + + gnms: dict[str, gnm_numpy.GNM] + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.gnms = {} + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS: + cls.gnms[version] = gnm_numpy.GNM.from_local( + gnm_numpy.GNMMajorVersion(version.removeprefix('v')), + gnm_numpy.GNMVariant.HEAD, + ) + + @parameterized.named_parameters(*[ + (version, version) + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS + ]) + def test_load_default_texture(self, version): + """Tests that load_texture returns the default texture.""" + gnm_np = self.gnms[version] + textures = render_common.load_texture(gnm_np) + self.assertIsInstance(textures, dict) + self.assertIn('skin', textures) + skin_tex = textures['skin'] + self.assertEqual(skin_tex.dtype, np.uint8) + self.assertEqual(skin_tex.shape[-1], 3) + + def test_load_custom_numpy_texture(self): + """Tests that load_texture returns the custom texture.""" + gnm_np = self.gnms[gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS[0]] + custom_tex = np.ones((32, 32, 3), dtype=np.float32) * 0.5 + textures = render_common.load_texture(gnm_np, texture=custom_tex) + self.assertIn('skin', textures) + self.assertEqual(textures['skin'].dtype, np.uint8) + self.assertEqual(textures['skin'].shape, (32, 32, 3)) + + +class TestProjectPointsForGnm(parameterized.TestCase): + """Tests for project_points_for_gnm.""" + + gnms: dict[str, gnm_numpy.GNM] + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.gnms = {} + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS: + cls.gnms[version] = gnm_numpy.GNM.from_local( + gnm_numpy.GNMMajorVersion(version.removeprefix('v')), + gnm_numpy.GNMVariant.HEAD, + ) + + @parameterized.named_parameters(*[ + (version, version) + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS + ]) + def test_project_points(self, version): + """Tests that project_points_for_gnm returns the correct points.""" + gnm_np = self.gnms[version] + vertices = gnm_np.template_vertex_positions + image_size = (240, 320) + projected = render_common.project_points_for_gnm( + gnm_np=gnm_np, + vertices=vertices, + image_size=image_size, + ) + self.assertEqual(projected.shape, (vertices.shape[0], 2)) + # Points should roughly fall within the image boundary. + self.assertTrue((projected[:, 0] >= -100).all()) + self.assertTrue((projected[:, 0] <= image_size[0] + 100).all()) + self.assertTrue((projected[:, 1] >= -100).all()) + self.assertTrue((projected[:, 1] <= image_size[1] + 100).all()) + + +class TestGetSpinWorldToCamera(parameterized.TestCase): + """Tests for get_spin_world_to_camera.""" + + gnms: dict[str, gnm_numpy.GNM] + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.gnms = {} + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS: + cls.gnms[version] = gnm_numpy.GNM.from_local( + gnm_numpy.GNMMajorVersion(version.removeprefix('v')), + gnm_numpy.GNMVariant.HEAD, + ) + + @parameterized.named_parameters(*[ + (version, version) + for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS + ]) + def test_spin_camera(self, version): + """Tests that get_spin_world_to_camera returns the correct matrices.""" + gnm_np = self.gnms[version] + vertices = gnm_np.template_vertex_positions + num_frames = 5 + spin_w2c = render_common.get_spin_world_to_camera( + gnm_np=gnm_np, + vertices=vertices, + num_frames=num_frames, + ) + self.assertEqual(spin_w2c.shape, (num_frames, 4, 4)) + + spin_period_w2c = render_common.get_spin_world_to_camera( + gnm_np=gnm_np, + vertices=vertices, + spin_period=num_frames, + ) + self.assertEqual(spin_period_w2c.shape, (num_frames, 4, 4)) + + +if __name__ == '__main__': + absltest.main() diff --git a/gnm/shape/visualization/render_gnm.py b/gnm/shape/visualization/render_gnm.py index e1bf1fa1..9f9aaa45 100644 --- a/gnm/shape/visualization/render_gnm.py +++ b/gnm/shape/visualization/render_gnm.py @@ -14,63 +14,33 @@ """Visualization utilities for GNM.""" -from collections.abc import Sequence -import functools - -from etils import epath from gnm.shape import gnm_numpy from gnm.shape.visualization import camera_conversions from gnm.shape.visualization import gnm_pyrender +from gnm.shape.visualization import render_common from gnm.shape.visualization import vertex_colors as vertex_colors_module -import imageio -import immutabledict import numpy as np import numpy.typing as npt -FloatArray = npt.NDArray[np.floating] - -ColorOrImage = npt.NDArray[np.uint8] | FloatArray | Sequence[float] | float - -_pkg = __package__ or 'gnm.shape.visualization' -_TEXTURES_DIR = epath.resource_path(_pkg).parent / 'data' / 'textures' -_EDGEFLOW_TEXTURE_BY_BODY_PART = immutabledict.immutabledict({ - gnm_numpy.GNMBodyPart.HEAD: str(_TEXTURES_DIR / 'edgeflow_bw_4k.png'), -}) - -# Default parameters for scene. -_DEFAULT_IMAGE_SIZE = (240, 320) -_DEFAULT_CAMERA_DISTANCE = 2.0 -_DEFAULT_NEAR = 0.01 -_DEFAULT_FAR = 100.0 -_DEFAULT_TARGET_FILL_FACTOR = 0.4 -_DEFAULT_BACKGROUND_COLOR = (0.95, 0.95, 0.95) - - -# Placeholder for lazy loading the default texture. -class _DefaultTexture: - """Placeholder for lazy loading the default texture.""" - - -DEFAULT_TEXTURE = _DefaultTexture() -Texture = FloatArray | _DefaultTexture | dict[str, FloatArray] | None - def render_gnm( gnm_np: gnm_numpy.GNM, - vertices: FloatArray | None = None, - world_to_camera: FloatArray | None = None, - camera_to_image: FloatArray | None = None, - image_size: tuple[int, int] = _DEFAULT_IMAGE_SIZE, + vertices: render_common.FloatArray | None = None, + world_to_camera: render_common.FloatArray | None = None, + camera_to_image: render_common.FloatArray | None = None, + image_size: tuple[int, int] = render_common.DEFAULT_IMAGE_SIZE, triangles: str | npt.NDArray[np.integer] = '~eye_exteriors', - texture: Texture = DEFAULT_TEXTURE, + texture: render_common.Texture = render_common.DEFAULT_TEXTURE, multisample_antialiasing: int = 2, - background_color: ColorOrImage = _DEFAULT_BACKGROUND_COLOR, + background_color: render_common.ColorOrImage = ( + render_common.DEFAULT_BACKGROUND_COLOR + ), alpha: float = 1.0, vertex_colors: npt.NDArray[np.floating] | None = None, multiple_gnms: bool = False, include_shading: bool = True, verbose: bool = False, -) -> FloatArray: +) -> render_common.FloatArray: """Render GNM meshes. Uses a pyrender backend to render GNM meshes. @@ -182,13 +152,13 @@ def render_gnm( vertices_for_cameras = vertices[..., 0, :, :] if world_to_camera is None: - world_to_camera = get_look_at_world_to_camera( + world_to_camera = render_common.get_look_at_world_to_camera( gnm_np, vertices_for_cameras, ) if camera_to_image is None: - camera_to_image = get_fill_factor_camera_to_image( + camera_to_image = render_common.get_fill_factor_camera_to_image( gnm_np, vertices_for_cameras, image_size=image_size ) @@ -201,8 +171,8 @@ def render_gnm( camera_to_image, width=image_size[0], height=image_size[1], - near=_DEFAULT_NEAR, - far=_DEFAULT_FAR, + near=render_common.DEFAULT_NEAR, + far=render_common.DEFAULT_FAR, ) ) @@ -216,9 +186,11 @@ def render_gnm( if background_color.dtype == np.uint8: background_color = background_color.astype(np.float32) / 255.0 - texture = _load_texture(gnm_np, texture) # pyrefly: ignore[bad-assignment] - textures = list(texture.values()) # pyrefly: ignore[missing-attribute] - texture_keys = texture.keys() # pyrefly: ignore[missing-attribute] + texture_dict = render_common.load_texture( + gnm_np, texture + ) + textures = list(texture_dict.values()) + texture_keys = texture_dict.keys() if not set(texture_keys).issubset(gnm_np.mesh_component_names): missing_parts = set(texture_keys) - set(gnm_np.mesh_component_names) raise ValueError( @@ -228,7 +200,7 @@ def render_gnm( # Find the maximum batch dimension that satisfies all batch-able arguments. try: - batch_dims = _get_batch_dim( + batch_dims = render_common.get_batch_dim( (vertices, 3), (vertex_colors, 3), (world_to_camera, 2), @@ -256,11 +228,9 @@ def batchify(arr, non_batch_dims): world_to_camera = batchify(world_to_camera, 2) camera_to_image = batchify(camera_to_image, 2) background_color = batchify(background_color, 3) - texture_dict = {} - for part in texture: # pyrefly: ignore[not-iterable] - x = texture[part] # pyrefly: ignore[bad-index, unsupported-operation] - texture_dict[part] = batchify(x, 3) - texture = texture_dict + batched_textures = { + part: batchify(x, 3) for part, x in texture_dict.items() + } renders = gnm_pyrender.render( vertices=vertices, @@ -268,7 +238,7 @@ def batchify(arr, non_batch_dims): world_to_camera=world_to_camera, camera_to_image=camera_to_image, image_size=image_size, - texture=texture, + texture=batched_textures, vertex_colors=vertex_colors, multisample_antialiasing=multisample_antialiasing, vertex_uvs=gnm_np.vertex_uvs, @@ -283,575 +253,3 @@ def batchify(arr, non_batch_dims): color = renders.reshape(*batch_dims, height, width, 3) return color - - -def project_points_for_gnm( - gnm_np: gnm_numpy.GNM, - points_world: np.ndarray | None = None, - vertices: np.ndarray | None = None, - world_to_camera: FloatArray | None = None, - camera_to_image: FloatArray | None = None, - image_size: tuple[int, int] = _DEFAULT_IMAGE_SIZE, - multiple_gnms: bool = False, - **kwargs, -) -> np.ndarray: - """Projects world points under the same conditions as a render_gnm call. - - Intended for identifying the per-frame positions of 3D points (e.g. GNM - joints) in the same reference space used by render_gnm. - - For a description of the other arguments, please see the docstring of - `render_gnm`. Any shading related arguments are ignored. - - Args: - gnm_np: The GNM model. - points_world: The world-space points to project, (..., P, 3). Defaults to - the vertex positions of the template GNM. - vertices: The GNM vertices in world space, (..., V, 3). If not provided, - will use the template vertices. Used for default camera setup. - world_to_camera: The world-to-camera transformation, (..., 4, 4). - camera_to_image: The camera-to-image transformation, (..., 4, 4). - image_size: The width and height of the rendered image in pixels: (W, H). - multiple_gnms: If True, vertices is expected to be shape (..., M, V, 3), and - we render M GNMs per image. Default cameras will be set relative to the - first GNM in the sequence. - **kwargs: Any additional arguments expected for render_gnm (ignored). - - Returns: - The projected points in image space, (..., P, 2). - """ - - del kwargs - - if points_world is None: - points_world = gnm_np.template_vertex_positions - points_world = points_world.astype(np.float32) - - if vertices is None: - vertices = gnm_np.template_vertex_positions - - if not multiple_gnms: - vertices = vertices[..., None, :, :] # Inject 'M' dimension. - elif (vertices_dim := vertices.ndim) < 3: - raise ValueError( - f'Called with {multiple_gnms=}, but vertices is only {vertices_dim}D.' - ) - - # Define default camera params based on the first GNM in the 'M' dimension. - vertices_for_cameras = vertices[..., 0, :, :] - - if world_to_camera is None: - world_to_camera = get_look_at_world_to_camera( - gnm_np, - vertices_for_cameras, - ) - - if camera_to_image is None: - camera_to_image = get_fill_factor_camera_to_image( - gnm_np, vertices_for_cameras, image_size=image_size - ) - - # Convert from OpenCV to OpenGL convention. - world_to_camera = camera_conversions.opencv_extrinsics_to_opengl( - world_to_camera - ) - camera_to_image = ( - camera_conversions.opencv_intrinsics_matrix_to_opengl_view_matrix( - camera_to_image, - width=image_size[0], - height=image_size[1], - near=_DEFAULT_NEAR, - far=_DEFAULT_FAR, - ) - ) - - # Find the maximum batch dimension that satisfies all batch-able arguments. - try: - batch_dims = _get_batch_dim( - (points_world, 2), - (vertices, 3), - (world_to_camera, 2), - (camera_to_image, 2), - ) - except ValueError as e: - raise ValueError( - f' Batch dimensions incompatible: points_world {points_world.shape},' - f' vertices {vertices.shape}, world_to_camera{{world_to_camera.shape}},' - ' camera_to_image {camera_to_image.shape}.' - ) from e - - def batchify(arr, non_batch_dims): - """Broadcast to batch dimensions, and flatten the batch dimensions.""" - arr = np.broadcast_to(arr, (*batch_dims, *arr.shape[-non_batch_dims:])) - return arr.reshape(int(np.prod(batch_dims)), *arr.shape[-non_batch_dims:]) - - points_world = batchify(points_world, 2) - world_to_camera = batchify(world_to_camera, 2) - camera_to_image = batchify(camera_to_image, 2) - - # Perform projection. - view_projection_matrix = camera_to_image @ world_to_camera - - homogenous_ones = np.ones( - (*points_world.shape[:-1], 1), dtype=points_world.dtype - ) - points_homogeneous = np.concatenate([points_world, homogenous_ones], axis=-1) - points_clip_space = ( - view_projection_matrix[:, None, :, :] @ points_homogeneous[..., None] - ) - points_clip_space = points_clip_space[..., 0] - - # Perform perspective division to get Normalized Device Coordinates (NDC). - points_ndc = points_clip_space[..., :3] / points_clip_space[..., [3]] - - # Convert NDC to image space. - # NDC range is [-1, 1]. Image space range is [0, image_size]. - width, height = image_size - points_image_space = ( - (points_ndc[..., :2] + 1.0) * 0.5 * np.array([width, height]) - ) - - # Flip +Y up renders to +Y down. - points_image_space[..., 1] = height - points_image_space[..., 1] - - return points_image_space.reshape(*batch_dims, *points_image_space.shape[-2:]) - - -def get_look_at_world_to_camera( - gnm_np: gnm_numpy.GNM, - vertices_world: np.ndarray, - azimuthal_angle: np.ndarray | float = 0.0, - polar_angle: np.ndarray | float = 0.0, - camera_distance: np.ndarray | float | None = _DEFAULT_CAMERA_DISTANCE, - share_camera: np.ndarray | bool = True, - y_up: np.ndarray | bool = False, - look_at_vertex_groups: Sequence[str] = ('hockey_mask',), - left_vertex_groups: Sequence[str] = ('ears', '&left'), - right_vertex_groups: Sequence[str] = ('ears', '&right'), - forward_vertex_groups: Sequence[str] = ('nose_region',), -) -> np.ndarray: - """Compute world-to-camera matrices for a 'look-at' transform. - - Returns matrices in OpenCV convention. - - Args: - gnm_np: The GNM model. - vertices_world: The GNM vertices in world space, (..., V, 3). - azimuthal_angle: The azimuthal angle of the camera in degrees, (..., 1). - polar_angle: The polar angle of the camera in degrees, (..., 1). - camera_distance: The distance of the camera from the head, (..., 1). - share_camera: Whether to use the first frame's vertices only for camera - generation, (..., 1). It is assumed that the first dimension of vertices - is the time dimension. - y_up: Whether to use the Y-up convention for the world space, (..., 1). - look_at_vertex_groups: The vertex groups to look at. - left_vertex_groups: The vertex groups to use for the left axis. - right_vertex_groups: The vertex groups to use for the right axis. - forward_vertex_groups: The vertex groups to use for the forward axis. - - Returns: - The world-to-camera matrices, (..., 4, 4). - """ - - batch_dims = vertices_world.shape[:-2] - - if not batch_dims: - # If there is no batch dimension, we don't need to share the camera. - share_camera = False - - azimuthal_angle = _adjust_scalar_shape(azimuthal_angle, batch_dims) - polar_angle = _adjust_scalar_shape(polar_angle, batch_dims) - camera_distance = _adjust_scalar_shape( - camera_distance, batch_dims # pyrefly: ignore[bad-argument-type] - ) - share_camera = _adjust_scalar_shape(share_camera, batch_dims) - y_up = _adjust_scalar_shape(y_up, batch_dims) - - first_frame_vertices = np.broadcast_to( - vertices_world[:1], vertices_world.shape - ) - - vertices_for_camera = np.where( - share_camera[..., None], first_frame_vertices, vertices_world - ) - - gnm_axes = _get_gnm_axes( - vertices_for_camera, - left_vertex_groups=left_vertex_groups, - right_vertex_groups=right_vertex_groups, - forward_vertex_groups=forward_vertex_groups, - gnm_np=gnm_np, - ) - camera_target = _vertex_group_mean( - vertices_for_camera, look_at_vertex_groups, gnm_np - ) - camera_location = camera_target + _get_camera_offset( - gnm_axes, - camera_distance, - azimuthal_angle, - polar_angle, - ) - - y_up_vector = np.array([0.0, 1.0, 0.0]) - up_vector = np.where(y_up, y_up_vector, gnm_axes[1]) - world_to_camera = _right_handed_look_at( - camera_location, camera_target, up_vector - ) - - world_to_camera = camera_conversions.opengl_extrinsics_to_opencv( - world_to_camera - ) - - return world_to_camera - - -def get_fill_factor_camera_to_image( - gnm_np: gnm_numpy.GNM, - vertices: np.ndarray, - target_fill_factor: np.ndarray | float = _DEFAULT_TARGET_FILL_FACTOR, - camera_distance: np.ndarray | float = _DEFAULT_CAMERA_DISTANCE, - image_size: tuple[int, int] = _DEFAULT_IMAGE_SIZE, - vertex_group_name: str = 'hockey_mask', -) -> np.ndarray: - """Compute the camera-to-image matrix for a 'fill factor' transform. - - The camera intrinsics are determined to fill the width of the given vertex - group with the image. Returns matrices in OpenCV convention. - - Args: - gnm_np: The GNM model. - vertices: The GNM vertices in world space, (..., V, 3). - target_fill_factor: The desired fill factor of the GNM mesh in the image, - (..., 1). - camera_distance: The distance of the camera from the head, (..., 1). - image_size: The width and height of the rendered image in pixels: (W, H). - vertex_group_name: The vertex group to use for the projection. - - Returns: - The camera-to-image matrix, (..., 4, 4). - """ - batch_dims = vertices.shape[:-2] - target_fill_factor = _adjust_scalar_shape(target_fill_factor, batch_dims) - camera_distance = _adjust_scalar_shape(camera_distance, batch_dims) - - width, height = image_size - - # Get the maximum width of the vertex group. - vertex_group_indices = gnm_np.vertex_group_indices(vertex_group_name) - vertex_group_points = vertices[..., vertex_group_indices, :] - mask_width = np.ptp(vertex_group_points, axis=-2)[..., :1] - dimension = mask_width / target_fill_factor - - vertical_field_of_view = np.atan2(dimension / 2.0, camera_distance) * 2 - aspect_ratio = _adjust_scalar_shape(width / height, batch_dims) - - near = _adjust_scalar_shape(_DEFAULT_NEAR, batch_dims) - far = _adjust_scalar_shape(_DEFAULT_FAR, batch_dims) - - camera_to_image = _right_handed_perspective( - vertical_field_of_view=vertical_field_of_view, - aspect_ratio=aspect_ratio, - near=near, - far=far, - ) - - camera_to_image = camera_conversions.opengl_intrinsics_to_opencv_matrix( - camera_to_image, width=width, height=height - ) - - return camera_to_image - - -def get_spin_world_to_camera( - gnm_np: gnm_numpy.GNM, - vertices: np.ndarray, - has_time_dimension: bool = False, - spin_period: int = 60, - spin_azimuth_limit: float = 20.0, - spin_polar_limit: float = 5.0, - **kwargs, -) -> np.ndarray: - """Compute world-to-camera matrices for a 'spin' transform. - - Args: - gnm_np: The GNM model. - vertices: The GNM vertices in world space, (..., V, 3). - has_time_dimension: Whether the first dimension of the vertices is time. If - false, will broadcast vertices over the spin period. - spin_period: The number of frames in a full spin. - spin_azimuth_limit: The maximum azimuthal angle of the camera in degrees. - spin_polar_limit: The maximum polar angle of the camera in degrees. - **kwargs: Additional arguments to pass to get_look_at_world_to_camera. - - Returns: - The world-to-camera matrices, (spin_period, ..., 4, 4) if has_time_dimension - is False, otherwise (..., 4, 4). - """ - - if has_time_dimension: - num_frames = vertices.shape[0] - else: - num_frames = spin_period - vertices = np.broadcast_to(vertices, (num_frames, *vertices.shape)) - - batch_dims = vertices.shape[:-2] - - # Spin frequency is determined by spin_period. - max_angle = np.pi * 2 * (num_frames / spin_period) - angles = np.linspace(-0, max_angle, num_frames) - - azimuthal_angle = np.sin(angles)[..., None] * spin_azimuth_limit - polar_angle = np.cos(angles)[..., None] * spin_polar_limit - - # Broadcast azimuthal and polar angles to the batch dimensions. - azimuthal_angle = np.broadcast_to(azimuthal_angle, (*batch_dims, 1)) - polar_angle = np.broadcast_to(polar_angle, (*batch_dims, 1)) - - return get_look_at_world_to_camera( - gnm_np, - vertices, - azimuthal_angle=azimuthal_angle, - polar_angle=polar_angle, - **kwargs, - ) - - -def _get_gnm_axes( - vertices: np.ndarray, - left_vertex_groups: Sequence[str], - right_vertex_groups: Sequence[str], - forward_vertex_groups: Sequence[str], - gnm_np: gnm_numpy.GNM, -) -> tuple[np.ndarray, np.ndarray, np.ndarray]: - """Determines the right, up, and forwards direction vectors of GNM.""" - - left = _vertex_group_mean(vertices, left_vertex_groups, gnm_np) - right = _vertex_group_mean(vertices, right_vertex_groups, gnm_np) - forward_point = _vertex_group_mean(vertices, forward_vertex_groups, gnm_np) - back_point = (left + right) / 2.0 - right = right - left - right = right / np.linalg.norm(right, axis=-1, keepdims=True) - forwards = forward_point - back_point - forwards = forwards / np.linalg.norm(forwards, axis=-1, keepdims=True) - up = np.linalg.cross(right, forwards) - return right, up, forwards - - -def _right_handed_look_at( - camera_position: np.ndarray, look_at: np.ndarray, up_vector: np.ndarray -): - """Builds a right handed look at view matrix.""" - z_axis = look_at - camera_position - z_axis /= np.linalg.norm(z_axis, axis=-1, keepdims=True) - - horizontal_axis = np.cross(z_axis, up_vector, axis=-1) - horizontal_axis /= np.linalg.norm(horizontal_axis, axis=-1, keepdims=True) - - vertical_axis = np.cross(horizontal_axis, z_axis, axis=-1) - - def _dot_last_axis(arr1, arr2): - return np.einsum('...i,...i->...', arr1, arr2)[..., None] - - batch_shape = horizontal_axis.shape[:-1] - zeros = np.zeros((*batch_shape, 3), dtype=horizontal_axis.dtype) - ones = np.ones((*batch_shape, 1), dtype=horizontal_axis.dtype) - - tx = -_dot_last_axis(horizontal_axis, camera_position) - ty = -_dot_last_axis(vertical_axis, camera_position) - tz = _dot_last_axis(z_axis, camera_position) - - row1 = np.concatenate([horizontal_axis, tx], axis=-1) - row2 = np.concatenate([vertical_axis, ty], axis=-1) - row3 = np.concatenate([-z_axis, tz], axis=-1) - row4 = np.concatenate([zeros, ones], axis=-1) - matrix = np.stack([row1, row2, row3, row4], axis=-2) - return matrix - - -def _get_camera_offset( - head_axes: tuple[np.ndarray, np.ndarray, np.ndarray], - camera_distance: np.ndarray, - azimuthal_angle: np.ndarray, - polar_angle: np.ndarray, -) -> np.ndarray: - """Determines how to place the camera so it points at the face. - - Args: - head_axes: Three-tuple representing normalised right, up, and forwards face - direction vectors, each (..., 3). - camera_distance: Camera distance from the face (meters), (..., 1). - azimuthal_angle: Places camera to the left or right, in degrees, (..., 1). - polar_angle: Places camera above or below the head, in degrees, (..., 1). - - Returns: - A batch of offset vectors to apply to camera targets (..., 3). - """ - right, up, forwards = head_axes - theta = np.deg2rad(90 + azimuthal_angle) - phi = np.deg2rad(90 + polar_angle) - x = camera_distance * np.sin(phi) * np.cos(theta) * right - y = camera_distance * np.sin(phi) * np.sin(theta) * forwards - z = camera_distance * np.cos(phi) * up - return x + y + z - - -def _right_handed_perspective( - vertical_field_of_view: np.ndarray, - aspect_ratio: np.ndarray, - near: np.ndarray, - far: np.ndarray, -) -> np.ndarray: - """Builds a right handed perspective projection matrix. - - Similar to tensorflow_graphics.rendering.camera.perspective.right_handed. - - Args: - vertical_field_of_view: The vertical field of view in radians, (..., 1). - aspect_ratio: The aspect ratio of the image, (..., 1). - near: The near clipping plane distance, (..., 1). - far: The far clipping plane distance, (..., 1). - - Returns: - The perspective projection matrix, (..., 4, 4) in OpenGL convention. - """ - - itan_half_vertical_field_of_view = 1.0 / np.tan(vertical_field_of_view * 0.5) - zero = np.zeros_like(itan_half_vertical_field_of_view) - one = np.ones_like(itan_half_vertical_field_of_view) - near_minus_far = near - far - - row1 = np.concatenate( - (itan_half_vertical_field_of_view / aspect_ratio, zero, zero, zero), - axis=-1, - ) - row2 = np.concatenate( - (zero, itan_half_vertical_field_of_view, zero, zero), axis=-1 - ) - row3 = np.concatenate( - ( - zero, - zero, - (far + near) / near_minus_far, - 2.0 * far * near / near_minus_far, - ), - axis=-1, - ) - row4 = np.concatenate((zero, zero, -one, zero), axis=-1) - - matrix = np.stack((row1, row2, row3, row4), axis=-2) - return matrix - - -@functools.cache -def _load_edgeflow_texture(texture_path: str | None) -> FloatArray | None: - if texture_path is None: - return None - with epath.Path(texture_path).open('rb') as f: - image = imageio.imread(f).astype(np.float32) - image = (image / np.iinfo(np.uint8).max) * 0.5 + 0.5 - image.flags.writeable = False - return image - - -def _vertex_group_mean( - vertices: np.ndarray, - group_names: Sequence[str], - gnm_np: gnm_numpy.GNM, -): - """Gets the average point of GNM vertex groups.""" - indices = gnm_np.vertex_group_indices(*group_names) - return vertices[..., indices, :].mean(axis=-2) - - -def _load_texture( - gnm_np: gnm_numpy.GNM, - texture: Texture = DEFAULT_TEXTURE, -) -> dict[str, npt.NDArray[np.uint8]]: - """Loads the texture as a (potentially batched) image. - - Args: - gnm_np: The GNM model. - texture: The texture to load. If _DEFAULT_TEXTURE, will load the edgeflow - texture. If an ndarray, will use the given texture for skin. If a dict, - will use the given texture for each part. If None, will use a white - (plain) texture. - - Returns: - The texture image, [0-255] uint8, (..., H, W, 3). - """ - texture_dict = {} - if texture is DEFAULT_TEXTURE: - edgeflow_path = _EDGEFLOW_TEXTURE_BY_BODY_PART.get(gnm_np.body_part) - if edgeflow_path is not None: - edgeflow = _load_edgeflow_texture(edgeflow_path) - texture_dict['skin'] = ( - edgeflow[..., None] # pyrefly: ignore[unsupported-operation] - ) - elif isinstance(texture, np.ndarray): - texture_dict['skin'] = texture - elif isinstance(texture, dict): - texture_dict = texture - - # Fill remaining parts with white texture. - for component in gnm_np.mesh_component_names: - if component not in texture_dict: - texture_dict[component] = np.ones((64, 64, 1)).astype(np.float32) - - def _to_3channel_uint8(texture_image: np.ndarray) -> npt.NDArray[np.uint8]: - texture_image = (texture_image * 255.0).astype(np.uint8) - if texture_image.shape[-1] == 1: - texture_image = np.repeat(texture_image, 3, axis=-1) - return texture_image - - texture_dict = { - part: _to_3channel_uint8(texture_dict[part]) for part in texture_dict - } - - return texture_dict - - -def _adjust_scalar_shape( - array: np.ndarray | float, batch_dims: Sequence[int] -) -> np.ndarray: - """Adjusts the shape of a scalar to match the batch dimensions.""" - - # If the array is a scalar or 1D array, add a final dimension of size 1. - array = np.array(array) - if array.ndim <= 1: - array = array[..., None] - - return np.broadcast_to(array, (*batch_dims, 1)) - - -def _get_batch_dim(*arrays: tuple[np.ndarray | None, int]) -> tuple[int, ...]: - """Finds the largest batch dimension that all arrays can be broadcast to. - - Requires that all arrays have batch dimensions that can be broadcast together. - - e.g. If given an array [B, C, 4] and an array [A, B, C, 3], return (A, B, C). - - Args: - *arrays: The arrays to expand. Tuples of (array, non_batch_dims), where - non_batch_dims is the number of rightmost dimensions of the array that are - not batch dimensions. If None, will be ignored. - - Returns: - The batch dimensions that all arrays can be safely broadcast to. - Raises: - ValueError: If the arrays cannot be broadcast together. - """ - batch_dims = [] - for array, non_batch_dims in arrays: - if array is None: - continue - batch_dims.append(array.shape[:-non_batch_dims]) - - max_batch_dims = max(batch_dims, key=len) - - # Verify that all arrays can be broadcast together. - for batch_dim in batch_dims: - np.broadcast_shapes( - batch_dim, max_batch_dims # pyrefly: ignore[bad-argument-type] - ) - - return max_batch_dims # pyrefly: ignore[bad-return] diff --git a/gnm/shape/visualization/render_gnm_test.py b/gnm/shape/visualization/render_gnm_test.py index 5ef8fd0a..df08fe7b 100644 --- a/gnm/shape/visualization/render_gnm_test.py +++ b/gnm/shape/visualization/render_gnm_test.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Tests for render_gnm.""" +"""Tests for rendering GNM meshes using the PyRender backend.""" # pylint: disable=protected-access @@ -25,7 +25,7 @@ from etils import epath from gnm.shape import gnm_numpy from gnm.shape.data.versions import gnm_test_catalog -from gnm.shape.visualization import camera_conversions +from gnm.shape.visualization import render_common from gnm.shape.visualization import render_gnm from gnm.shape.visualization import vertex_colors as vertex_colors_module import mediapy as media @@ -141,7 +141,7 @@ def test_spin_period(self, version, spin_period: int): """Tests we can render spins of different length.""" gnm_np = self.gnms[version] - world_to_camera = render_gnm.get_spin_world_to_camera( + world_to_camera = render_common.get_spin_world_to_camera( gnm_np=gnm_np, vertices=gnm_np.template_vertex_positions, spin_period=spin_period, @@ -173,7 +173,7 @@ def test_spin_period_with_time_dimension(self, version): parameters = _get_random_parameters((num_frames,), gnm_np) vertices = gnm_np(**parameters) - world_to_camera = render_gnm.get_spin_world_to_camera( + world_to_camera = render_common.get_spin_world_to_camera( gnm_np=gnm_np, vertices=vertices, has_time_dimension=True, @@ -486,7 +486,7 @@ def test_per_part_texture(self, version): ) eye_indices = gnm_np.vertex_group_indices('eye_interiors') - eye_image_points = render_gnm.project_points_for_gnm( + eye_image_points = render_common.project_points_for_gnm( points_world=vertices[eye_indices], vertices=vertices, gnm_np=gnm_np, @@ -690,26 +690,6 @@ def test_error_on_batch_mismatch(self, version): ) -class TestGetBatchDim(parameterized.TestCase): - """Tests for _get_batch_dim helper.""" - - def test_get_batch_dim(self): - a, b, c = 1, 2, 3 - array_a = np.zeros((a, b, c, 10)) - array_b = None - array_c = np.zeros((b, c, 5, 5)) - batch_dims = render_gnm._get_batch_dim( - (array_a, 1), (array_b, 2), (array_c, 2) - ) - self.assertEqual(batch_dims, (a, b, c)) - - def test_get_batch_dim_raises_error(self): - array_a = np.zeros((1, 2, 3, 10)) - array_b = np.zeros((4, 5, 6, 5)) - with self.assertRaises(ValueError): - render_gnm._get_batch_dim((array_a, 1), (array_b, 1)) - - class TestProjectPointsForGNM(parameterized.TestCase): """Tests projection of points for GNM.""" @@ -753,7 +733,7 @@ def test_project_points_default_render(self, version): image = render_gnm.render_gnm(gnm_np, **self.rendering_kwargs) # Project face joints under the same camera setup. - joints_image = render_gnm.project_points_for_gnm( + joints_image = render_common.project_points_for_gnm( gnm_np=gnm_np, points_world=gnm_np.template_joint_positions, **self.rendering_kwargs, @@ -763,7 +743,7 @@ def test_project_points_default_render(self, version): points_world = np.array( [[0, gnm_np.template_vertex_positions[:, 1].max() + 0.05, 0]] ) - external_point_image = render_gnm.project_points_for_gnm( + external_point_image = render_common.project_points_for_gnm( gnm_np=gnm_np, points_world=points_world, **self.rendering_kwargs, @@ -796,7 +776,7 @@ def test_project_points_spin(self, version): """Tests projection of points for GNM in a spin.""" gnm_np = self.gnms[version] spin_period = 30 - world_to_camera = render_gnm.get_spin_world_to_camera( + world_to_camera = render_common.get_spin_world_to_camera( gnm_np=gnm_np, vertices=gnm_np.template_vertex_positions, spin_period=spin_period, @@ -807,7 +787,7 @@ def test_project_points_spin(self, version): ) # Project face joints under the same camera setup. - joints_image = render_gnm.project_points_for_gnm( + joints_image = render_common.project_points_for_gnm( gnm_np=gnm_np, points_world=gnm_np.template_joint_positions, world_to_camera=world_to_camera, @@ -840,155 +820,5 @@ def test_project_points_spin(self, version): ) -class TestGetLookAtWorldToCamera(parameterized.TestCase): - """Tests for get_look_at_world_to_camera.""" - - gnms: dict[str, gnm_numpy.GNM] - - @classmethod - def setUpClass(cls): - super().setUpClass() - cls.gnms = {} - for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS: - cls.gnms[version] = gnm_numpy.GNM.from_local( - gnm_numpy.GNMMajorVersion(version.removeprefix('v')), - gnm_numpy.GNMVariant.HEAD, - ) - - @parameterized.named_parameters(*[ - (version, version) - for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS - ]) - def test_basic(self, version): - gnm_np = self.gnms[version] - vertices = gnm_np.template_vertex_positions - world_to_camera_opencv = render_gnm.get_look_at_world_to_camera( - gnm_np=gnm_np, - vertices_world=vertices, - ) - world_to_camera_opengl = camera_conversions.opencv_extrinsics_to_opengl( - world_to_camera_opencv - ) - self.assertEqual(world_to_camera_opengl.shape, (4, 4)) - - with self.subTest('Approximately identity rotation.'): - np.testing.assert_allclose( - world_to_camera_opengl[:3, :3], np.eye(3), atol=0.02 - ) - - with self.subTest('Translated from hockey mask.'): - hockey_mask_indices = gnm_np.vertex_group_indices('hockey_mask') - hockey_mask_z = vertices[hockey_mask_indices, 2].mean() - self.assertAlmostEqual( - -world_to_camera_opengl[2, 3], - hockey_mask_z + render_gnm._DEFAULT_CAMERA_DISTANCE, - delta=0.01, - ) - - @parameterized.named_parameters(*[ - (version, version) - for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS - ]) - def test_share_camera_no_batch(self, version): - gnm_np = self.gnms[version] - parameters = _get_random_parameters((), gnm_np) - vertices = gnm_np(**parameters) - world_to_camera_no_share = render_gnm.get_look_at_world_to_camera( - gnm_np=gnm_np, - vertices_world=vertices, - share_camera=False, - ) - - world_to_camera_share = render_gnm.get_look_at_world_to_camera( - gnm_np=gnm_np, - vertices_world=vertices, - share_camera=True, - ) - - np.testing.assert_allclose(world_to_camera_no_share, world_to_camera_share) - - @parameterized.named_parameters(*[ - (version, version) - for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS - ]) - def test_share_camera(self, version): - gnm_np = self.gnms[version] - parameters = _get_random_parameters((10,), gnm_np) - vertices = gnm_np(**parameters) - - world_to_camera_0 = render_gnm.get_look_at_world_to_camera( - gnm_np=gnm_np, vertices_world=vertices[0] - ) - - world_to_camera_share = render_gnm.get_look_at_world_to_camera( - gnm_np=gnm_np, - vertices_world=vertices, - share_camera=True, - ) - - world_to_camera_no_share = render_gnm.get_look_at_world_to_camera( - gnm_np=gnm_np, - vertices_world=vertices, - share_camera=False, - ) - - with self.subTest('Shared camera matches first camera.'): - np.testing.assert_allclose( - world_to_camera_share, - np.broadcast_to(world_to_camera_0, world_to_camera_share.shape), - ) - - with self.subTest('Unshared cameras are all different.'): - self.assertFalse( - np.allclose( - world_to_camera_no_share, - np.broadcast_to( - world_to_camera_0, world_to_camera_no_share.shape - ), - ) - ) - - -class TestGetFillFactorCameraToImage(parameterized.TestCase): - """Tests for get_fill_factor_camera_to_image.""" - - gnms: dict[str, gnm_numpy.GNM] - - @classmethod - def setUpClass(cls): - super().setUpClass() - cls.gnms = {} - for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS: - cls.gnms[version] = gnm_numpy.GNM.from_local( - gnm_numpy.GNMMajorVersion(version.removeprefix('v')), - gnm_numpy.GNMVariant.HEAD, - ) - - @parameterized.named_parameters(*[ - (version, version) - for version in gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS - ]) - def test_basic(self, version): - gnm_np = self.gnms[version] - vertices = gnm_np.template_vertex_positions - camera_to_image_opencv = render_gnm.get_fill_factor_camera_to_image( - gnm_np=gnm_np, - vertices=vertices, - ) - camera_to_image_opengl = ( - camera_conversions.opencv_intrinsics_matrix_to_opengl_view_matrix( - camera_to_image_opencv, - width=320, - height=240, - near=0.1, - far=100.0, - ) - ) - self.assertEqual(camera_to_image_opengl.shape, (4, 4)) - with self.subTest('Lower triangular is all zeros.'): - submatrix = camera_to_image_opengl[:3, :3] - self.assertTrue((np.tril(submatrix, k=-1) == 0.0).all()) - - if __name__ == '__main__': absltest.main()