diff --git a/gnm/shape/data/versions/gnm_specs.py b/gnm/shape/data/versions/gnm_specs.py index c8106600..75568ba1 100644 --- a/gnm/shape/data/versions/gnm_specs.py +++ b/gnm/shape/data/versions/gnm_specs.py @@ -32,3 +32,8 @@ class GNMBodyPart(enum.StrEnum): GNM_VARIANT_TO_BODY_PART_MAP = { GNMVariant.HEAD: GNMBodyPart.HEAD, } + +class GNMRemoteSource(enum.StrEnum): + HTTP = 'http' + HUGGING_FACE = 'huggingface' + KAGGLE = 'kaggle' diff --git a/gnm/shape/data/versions/v3_0/gnm_head.npz b/gnm/shape/data/versions/v3_0/gnm_head.npz index 0b3a3f33..365e6f2b 100644 Binary files a/gnm/shape/data/versions/v3_0/gnm_head.npz and b/gnm/shape/data/versions/v3_0/gnm_head.npz differ diff --git a/gnm/shape/demos/gnm_head_demo.ipynb b/gnm/shape/demos/gnm_head_demo.ipynb index fd861e27..a673be9c 100644 --- a/gnm/shape/demos/gnm_head_demo.ipynb +++ b/gnm/shape/demos/gnm_head_demo.ipynb @@ -15,6 +15,17 @@ ], "metadata": {} }, + { + "id": "74eef724", + "cell_type": "code", + "source": [ + "# @title Install GNM and dependencies\n", + "!pip install -q git+https://github.com/google/GNM.git#subdirectory=gnm/shape\n" + ], + "metadata": {}, + "execution_count": null, + "outputs": [] + }, { "id": "2fdljmx_w4hW", "cell_type": "code", @@ -44,10 +55,25 @@ "source": [ "# @title Instantiate GNM\n", "\n", - "gnm = gnm_numpy.GNM.from_local(\n", + "# 1. Direct HTTPS CDN stream (Default, zero extra dependencies required)\n", + "gnm = gnm_numpy.GNM.from_remote(\n", " version=gnm_numpy.GNMMajorVersion.V3,\n", - " variant=gnm_numpy.GNMVariant.HEAD\n", - ")\n" + " variant=gnm_numpy.GNMVariant.HEAD,\n", + ")\n", + "\n", + "# 2. Or from Hugging Face Hub (requires huggingface_hub: pip install huggingface_hub)\n", + "# gnm = gnm_numpy.GNM.from_remote(\n", + "# version=gnm_numpy.GNMMajorVersion.V3,\n", + "# variant=gnm_numpy.GNMVariant.HEAD,\n", + "# source=gnm_numpy.GNMRemoteSource.HUGGING_FACE,\n", + "# )\n", + "\n", + "# 3. Or from Kaggle Models (requires kagglehub: pip install kagglehub)\n", + "# gnm = gnm_numpy.GNM.from_remote(\n", + "# version=gnm_numpy.GNMMajorVersion.V3,\n", + "# variant=gnm_numpy.GNMVariant.HEAD,\n", + "# source=gnm_numpy.GNMRemoteSource.KAGGLE,\n", + "# )\n" ], "metadata": {}, "execution_count": null, diff --git a/gnm/shape/demos/semantic_gnm_demo.ipynb b/gnm/shape/demos/semantic_gnm_demo.ipynb index 0e9e252a..331a1240 100644 --- a/gnm/shape/demos/semantic_gnm_demo.ipynb +++ b/gnm/shape/demos/semantic_gnm_demo.ipynb @@ -49,10 +49,26 @@ "# @title Instantiate Models\n", "# @test {\"skip\": true}\n", "\n", - "gnm = gnm_numpy.GNM.from_local(\n", + "# 1. Direct HTTPS CDN stream (Default, zero extra dependencies required)\n", + "gnm = gnm_numpy.GNM.from_remote(\n", " version=gnm_numpy.GNMMajorVersion.V3,\n", - " variant=gnm_numpy.GNMVariant.HEAD\n", + " variant=gnm_numpy.GNMVariant.HEAD,\n", ")\n", + "\n", + "# 2. Or from Hugging Face Hub (requires huggingface_hub: pip install huggingface_hub)\n", + "# gnm = gnm_numpy.GNM.from_remote(\n", + "# version=gnm_numpy.GNMMajorVersion.V3,\n", + "# variant=gnm_numpy.GNMVariant.HEAD,\n", + "# source=gnm_numpy.GNMRemoteSource.HUGGING_FACE,\n", + "# )\n", + "\n", + "# 3. Or from Kaggle Models (requires kagglehub: pip install kagglehub)\n", + "# gnm = gnm_numpy.GNM.from_remote(\n", + "# version=gnm_numpy.GNMMajorVersion.V3,\n", + "# variant=gnm_numpy.GNMVariant.HEAD,\n", + "# source=gnm_numpy.GNMRemoteSource.KAGGLE,\n", + "# )\n", + "\n", "expr_sampler = semantic_sampler.ExpressionSampler()\n", "id_sampler = semantic_sampler.IdentitySampler()\n" ], diff --git a/gnm/shape/gnm_base.py b/gnm/shape/gnm_base.py index b4a1f820..e5cafc43 100644 --- a/gnm/shape/gnm_base.py +++ b/gnm/shape/gnm_base.py @@ -21,6 +21,7 @@ import dataclasses from typing import Any, Self +from etils import epath from gnm.shape import gnm_data_loader from gnm.shape.data.versions import gnm_specs @@ -42,6 +43,38 @@ def from_local( data_dict = gnm_data_loader.load_model_from_runfile(version, variant) return cls._from_model_data(data_dict) # pyrefly: ignore[bad-return] + @classmethod + def from_remote( + cls, + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + source: gnm_specs.GNMRemoteSource = gnm_specs.GNMRemoteSource.HTTP, + *, + cache_dir: epath.PathLike | None = None, + force_download: bool = False, + ) -> Self: + """Creates a GNM instance from a remote repository. + + Args: + version: GNM major version. + variant: GNM model variant. + source: Remote repository source (HTTP, Hugging Face, or Kaggle). + cache_dir: Optional custom directory (Path or str) to cache downloaded + models. + force_download: If True, forces redownload even if cached locally. + + Returns: + A GNM instance loaded with the model weights. + """ + data_dict = gnm_data_loader.load_model_from_remote( + version=version, + variant=variant, + source=source, + cache_dir=cache_dir, + force_download=force_download, + ) + return cls._from_model_data(data_dict) + @classmethod def from_gnm(cls, gnm: GNMBase) -> Self: """Creates a GNM instance from another GNM instance.""" diff --git a/gnm/shape/gnm_base_test.py b/gnm/shape/gnm_base_test.py index 7a78edd3..eb1dfc88 100644 --- a/gnm/shape/gnm_base_test.py +++ b/gnm/shape/gnm_base_test.py @@ -12,12 +12,13 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Tests for gnm_base.""" +"""Unit tests for GNMBase factory methods, properties, and serialization.""" from __future__ import annotations from collections.abc import Mapping from typing import Any +from unittest import mock from absl.testing import absltest from gnm.shape import gnm_base @@ -50,7 +51,7 @@ def __init__( self.variant = variant def to_numpy_data_dict(self) -> dict[str, Any]: - return {"dummy": 1} + return {'dummy': 1} @classmethod def _from_model_data( @@ -83,6 +84,72 @@ def test_from_gnm(self): self.assertEqual(new_gnm.version, _TEST_FULL_VERSION) self.assertEqual(new_gnm.variant, _TEST_VARIANT) - -if __name__ == "__main__": + def test_from_local(self): + with mock.patch.object( + gnm_data_loader, + 'load_model_from_runfile', + return_value={'dummy': 1}, + ) as mock_load: + new_gnm = DummyGNM.from_local(_TEST_MAJOR_VERSION, _TEST_VARIANT) + self.assertIsInstance(new_gnm, DummyGNM) + mock_load.assert_called_once_with(_TEST_MAJOR_VERSION, _TEST_VARIANT) + + def test_from_remote_default_http(self): + with mock.patch.object( + gnm_data_loader, + 'load_model_from_remote', + return_value={'dummy': 1}, + ) as mock_load: + new_gnm = DummyGNM.from_remote(_TEST_MAJOR_VERSION, _TEST_VARIANT) + self.assertIsInstance(new_gnm, DummyGNM) + mock_load.assert_called_once_with( + version=_TEST_MAJOR_VERSION, + variant=_TEST_VARIANT, + source=gnm_specs.GNMRemoteSource.HTTP, + cache_dir=None, + force_download=False, + ) + + def test_from_remote_huggingface(self): + with mock.patch.object( + gnm_data_loader, + 'load_model_from_remote', + return_value={'dummy': 1}, + ) as mock_load: + new_gnm = DummyGNM.from_remote( + _TEST_MAJOR_VERSION, + _TEST_VARIANT, + source=gnm_specs.GNMRemoteSource.HUGGING_FACE, + ) + self.assertIsInstance(new_gnm, DummyGNM) + mock_load.assert_called_once_with( + version=_TEST_MAJOR_VERSION, + variant=_TEST_VARIANT, + source=gnm_specs.GNMRemoteSource.HUGGING_FACE, + cache_dir=None, + force_download=False, + ) + + def test_from_remote_kaggle(self): + with mock.patch.object( + gnm_data_loader, + 'load_model_from_remote', + return_value={'dummy': 1}, + ) as mock_load: + new_gnm = DummyGNM.from_remote( + _TEST_MAJOR_VERSION, + _TEST_VARIANT, + source=gnm_specs.GNMRemoteSource.KAGGLE, + ) + self.assertIsInstance(new_gnm, DummyGNM) + mock_load.assert_called_once_with( + version=_TEST_MAJOR_VERSION, + variant=_TEST_VARIANT, + source=gnm_specs.GNMRemoteSource.KAGGLE, + cache_dir=None, + force_download=False, + ) + + +if __name__ == '__main__': absltest.main() diff --git a/gnm/shape/gnm_data_loader.py b/gnm/shape/gnm_data_loader.py index b260c58f..f00429a7 100644 --- a/gnm/shape/gnm_data_loader.py +++ b/gnm/shape/gnm_data_loader.py @@ -14,7 +14,6 @@ """GNM data loader.""" -# from collections.abc import Mapping, Sequence from collections.abc import Sequence import functools from typing import Any @@ -39,17 +38,6 @@ class GNMModelDataNotLinkedError(Exception): pass -def _get_model_path_from_version_and_variant( - version: gnm_specs.GNMMajorVersion, - variant: gnm_specs.GNMVariant, -) -> epath.Path: - """Returns the GNM model runfiles path for given variant and version.""" - version_value = major_to_newest_full_version(version).value.replace('.', '_') - version_dir_name = f'v{version_value}' - model_file_name = f'{_VARIANT_TO_MODEL_FILE_NAME_MAP[variant]}.npz' - return _MODELS_VERSIONS_DIR / version_dir_name / model_file_name - - def major_to_newest_full_version( major: gnm_specs.GNMMajorVersion, ) -> gnm_specs.GNMVersion: @@ -67,6 +55,31 @@ def full_version_to_major( return gnm_specs.GNMMajorVersion(version.value.split('.')[0]) +def _get_version_dir_name(version: gnm_specs.GNMMajorVersion) -> str: + """Returns directory name for version (e.g. 'v3_0').""" + version_value = major_to_newest_full_version(version).value.replace('.', '_') + return f'v{version_value}' + + +def _get_model_filename( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, +) -> str: + """Returns filename for model variant (e.g. 'gnm_head.npz').""" + del version + return f'{_VARIANT_TO_MODEL_FILE_NAME_MAP[variant]}.npz' + + +def _get_model_path_from_version_and_variant( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, +) -> epath.Path: + """Returns the GNM model runfiles path for given variant and version.""" + version_dir_name = _get_version_dir_name(version) + model_file_name = _get_model_filename(version, variant) + return _MODELS_VERSIONS_DIR / version_dir_name / model_file_name + + @functools.lru_cache def load_model_from_runfile( version: gnm_specs.GNMMajorVersion, variant: gnm_specs.GNMVariant @@ -80,20 +93,47 @@ def load_model_from_runfile( variant, model_file, ) + return _load_model_dict_from_file(model_file, version, variant) + + +def _load_model_dict_from_file( + model_file: epath.Path, + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, +) -> dict[str, Any]: + """Loads and standardizes model dict from a local file path.""" with model_file.open('rb') as f: data_dict = dict(np.load(f)) + del version, variant + # Validate the data. valid, missing, extra = _validate_gnm_data(data_dict) if not valid: raise ValueError( - f'Validation failed for version {version}, variant {variant}.' + f'Validation failed for model from {model_file}.' f' Missing: {missing}, Extra: {extra}' ) return _standardize_gnm_data_types(data_dict) +def get_default_gnm_cache_dir() -> epath.Path: + """Returns the default directory for caching downloaded GNM models.""" + from gnm.shape.oss_data_loaders import oss_data_loaders # pylint: disable=g-import-not-at-top,import-outside-toplevel + return oss_data_loaders.get_default_gnm_cache_dir() + + +def load_model_from_remote( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + **kwargs: Any, +) -> dict[str, Any]: + """Loads GNM model data from a remote source.""" + from gnm.shape.oss_data_loaders import oss_data_loaders # pylint: disable=g-import-not-at-top,import-outside-toplevel + return oss_data_loaders.load_model_from_remote(version, variant, **kwargs) + + def _validate_gnm_data( data: dict[str, Any], ) -> tuple[bool, Sequence[str], Sequence[str]]: @@ -109,9 +149,9 @@ def _validate_gnm_data( A tuple of (bool, Sequence[str], Sequence[str]) indicating if the data dict has exactly the expected fields, the missing fields and the extra fields. """ - expected_fields = gnm_data_schema.GNM_DATA_ATTRIBUTES - missing_fields = list(set(expected_fields) - set(data.keys())) - extra_fields = list(set(data.keys()) - set(expected_fields)) + expected_fields = set(gnm_data_schema.GNM_DATA_ATTRIBUTES) + missing_fields = list(expected_fields - set(data.keys())) + extra_fields = list(set(data.keys()) - expected_fields) return not missing_fields and not extra_fields, missing_fields, extra_fields diff --git a/gnm/shape/gnm_data_loader_test.py b/gnm/shape/gnm_data_loader_test.py index e51a3e8b..adb23cbb 100644 --- a/gnm/shape/gnm_data_loader_test.py +++ b/gnm/shape/gnm_data_loader_test.py @@ -12,8 +12,16 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Tests for gnm_data_loader.""" +"""Unit tests verifying GNM model data loaders. +Tests loaders across runfiles and TFHub. +""" + +# pylint: disable=protected-access + +import io +import os +from typing import Any from unittest import mock from absl.testing import absltest @@ -22,11 +30,42 @@ from gnm.shape import gnm_data_loader from gnm.shape.data.versions import gnm_specs from gnm.shape.data.versions import gnm_test_catalog +from gnm.shape.oss_data_loaders import oss_data_loaders +import numpy as np _MAINTAINED_MAJOR_GNM_VERSIONS = gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS _MAJOR_VERSION_TO_VARIANTS_MAP = gnm_test_catalog.MAJOR_VERSION_TO_VARIANTS_MAP +def _get_dummy_gnm_data_dict() -> dict[str, Any]: + """Returns a dummy GNM data dictionary.""" + return { + 'version': '3.0', + 'variant': 'head', + 'template_vertex_positions': np.zeros((1, 3)), + 'template_joint_positions': np.zeros((1, 3)), + 'vertex_identity_basis': np.zeros((1, 1, 3)), + 'joint_identity_basis': np.zeros((1, 1, 3)), + 'expression_basis': np.zeros((1, 1, 3)), + 'identity_names': ['id1'], + 'joint_names': ['joint1'], + 'expression_names': ['exp1'], + 'joint_parent_indices': np.array([0]), + 'skinning_weights': np.zeros((1, 1)), + 'quads': np.zeros((1, 4)), + 'triangles': np.zeros((1, 3)), + 'quad_uvs': np.zeros((1, 4, 2)), + 'triangle_uvs': np.zeros((1, 3, 2)), + 'mesh_component_names': ['part1'], + 'mirror_indices': np.array([0]), + 'joint_regressor': np.zeros((1, 1)), + 'pose_correctives_regressor': np.zeros((9, 3)), + 'bone_aligned_template_joint_orientations': np.zeros((1, 3, 3)), + 'vertex_groups': np.zeros((1, 1)), + 'vertex_group_names': ['group1'], + } + + class GNMDataTest(parameterized.TestCase): def test_print_gnm_major_versions(self): @@ -45,6 +84,10 @@ def test_print_gnm_versions(self): class GNMModelLoadingTest(parameterized.TestCase): """Tests for loading GNM model files.""" + def setUp(self): + super().setUp() + gnm_data_loader.load_model_from_runfile.cache_clear() + @parameterized.product( version=_MAINTAINED_MAJOR_GNM_VERSIONS, variant=gnm_test_catalog.ALL_VARIANTS, @@ -73,5 +116,160 @@ def test_load_model_from_runfile_fails_when_file_not_found(self): ) +class GNMRemoteModelLoadingTest(parameterized.TestCase): + """Tests for remote model loading and caching in gnm_data_loader.""" + + def setUp(self): + super().setUp() + self.temp_dir = epath.Path(self.create_tempdir().full_path) + self.dummy_gnm_data_dict = _get_dummy_gnm_data_dict() + + # Save a dummy npz in temp_dir. + buffer = io.BytesIO() + np.savez_compressed(buffer, **self.dummy_gnm_data_dict) + self.dummy_npz_bytes = buffer.getvalue() + + # Dynamically determine the latest maintained version and an available + # variant. + version_key = gnm_test_catalog.MAINTAINED_MAJOR_VERSIONS[-1] + self.test_version = gnm_specs.GNMMajorVersion( + version_key.removeprefix('v') + ) + variants = gnm_test_catalog.MAJOR_VERSION_TO_VARIANTS_MAP[version_key] + if gnm_specs.GNMVariant.HEAD.value in variants: + self.test_variant = gnm_specs.GNMVariant.HEAD + else: + self.test_variant = gnm_specs.GNMVariant(variants[0]) + version_dir_name = gnm_data_loader._get_version_dir_name( + self.test_version + ) + model_file_name = gnm_data_loader._get_model_filename( + self.test_version, self.test_variant + ) + self.dest_cache_file = self.temp_dir / version_dir_name / model_file_name + + def test_get_default_gnm_cache_dir(self): + with mock.patch.dict(os.environ, {'GNM_CACHE_DIR': '/custom/gnm/cache'}): + self.assertEqual( + gnm_data_loader.get_default_gnm_cache_dir(), + epath.Path('/custom/gnm/cache'), + ) + + with mock.patch.dict( + os.environ, + {'XDG_CACHE_HOME': '/custom/xdg/cache'}, + clear=True, + ): + self.assertEqual( + gnm_data_loader.get_default_gnm_cache_dir(), + epath.Path('/custom/xdg/cache/gnm/models'), + ) + + def test_load_model_from_remote(self): + def _fake_download(url, dest): + del url + dest.parent.mkdir(parents=True, exist_ok=True) + dest.write_bytes(self.dummy_npz_bytes) + return dest + + with mock.patch.object( + oss_data_loaders, '_download_file', side_effect=_fake_download + ) as mock_download: + # First load: triggers download + data1 = gnm_data_loader.load_model_from_remote( + self.test_version, + self.test_variant, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data1, dict) + self.assertEqual(mock_download.call_count, 1) + self.assertTrue(self.dest_cache_file.exists()) + + # Second load: uses cached file directly, does not re-download + data2 = gnm_data_loader.load_model_from_remote( + self.test_version, + self.test_variant, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data2, dict) + self.assertEqual(mock_download.call_count, 1) + + # Force download: re-downloads even if cached + data3 = gnm_data_loader.load_model_from_remote( + self.test_version, + self.test_variant, + cache_dir=self.temp_dir, + force_download=True, + ) + self.assertIsInstance(data3, dict) + self.assertEqual(mock_download.call_count, 2) + + def test_load_model_from_remote_with_str_cache_dir(self): + def _fake_download(url, dest): + del url + dest.parent.mkdir(parents=True, exist_ok=True) + dest.write_bytes(self.dummy_npz_bytes) + return dest + + with mock.patch.object( + oss_data_loaders, '_download_file', side_effect=_fake_download + ): + data = gnm_data_loader.load_model_from_remote( + self.test_version, + self.test_variant, + cache_dir=str(self.temp_dir), + ) + self.assertIsInstance(data, dict) + self.assertTrue(self.dest_cache_file.exists()) + + def test_load_model_from_remote_huggingface(self): + dest_file = self.dest_cache_file + dest_file.parent.mkdir(parents=True, exist_ok=True) + dest_file.write_bytes(self.dummy_npz_bytes) + + with mock.patch.object( + oss_data_loaders, + '_resolve_huggingface_model_file', + return_value=dest_file, + ) as mock_resolve: + data = gnm_data_loader.load_model_from_remote( + self.test_version, + self.test_variant, + source=gnm_specs.GNMRemoteSource.HUGGING_FACE, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data, dict) + mock_resolve.assert_called_once_with( + self.test_version, + self.test_variant, + self.temp_dir, + False, + ) + + def test_load_model_from_remote_kaggle(self): + dest_file = self.dest_cache_file + dest_file.parent.mkdir(parents=True, exist_ok=True) + dest_file.write_bytes(self.dummy_npz_bytes) + + with mock.patch.object( + oss_data_loaders, + '_resolve_kaggle_model_file', + return_value=dest_file, + ) as mock_resolve: + data = gnm_data_loader.load_model_from_remote( + self.test_version, + self.test_variant, + source=gnm_specs.GNMRemoteSource.KAGGLE, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data, dict) + mock_resolve.assert_called_once_with( + self.test_version, + self.test_variant, + self.temp_dir, + False, + ) + + if __name__ == '__main__': absltest.main() diff --git a/gnm/shape/gnm_jax.py b/gnm/shape/gnm_jax.py index fb567838..fed94dad 100644 --- a/gnm/shape/gnm_jax.py +++ b/gnm/shape/gnm_jax.py @@ -54,6 +54,7 @@ GNMVariant = gnm_specs.GNMVariant GNMBodyPart = gnm_specs.GNMBodyPart GNMLandmarksType = gnm_landmarks.GNMLandmarksType +GNMRemoteSource = gnm_specs.GNMRemoteSource @dataclasses.dataclass(frozen=False, kw_only=True, init=False) diff --git a/gnm/shape/gnm_numpy.py b/gnm/shape/gnm_numpy.py index 3c5609a3..ede1e47a 100644 --- a/gnm/shape/gnm_numpy.py +++ b/gnm/shape/gnm_numpy.py @@ -49,6 +49,7 @@ GNMVariant = gnm_specs.GNMVariant GNMBodyPart = gnm_specs.GNMBodyPart GNMLandmarksType = gnm_landmarks.GNMLandmarksType +GNMRemoteSource = gnm_specs.GNMRemoteSource _rotation_matrix = gnm_common.axis_angle_to_rotation_matrix diff --git a/gnm/shape/gnm_pytorch.py b/gnm/shape/gnm_pytorch.py index 76cde071..3aa6a9c7 100644 --- a/gnm/shape/gnm_pytorch.py +++ b/gnm/shape/gnm_pytorch.py @@ -49,6 +49,7 @@ GNMVariant = gnm_specs.GNMVariant GNMBodyPart = gnm_specs.GNMBodyPart GNMLandmarksType = gnm_landmarks.GNMLandmarksType +GNMRemoteSource = gnm_specs.GNMRemoteSource @dataclasses.dataclass(frozen=False, kw_only=True, init=False) diff --git a/gnm/shape/gnm_tensorflow.py b/gnm/shape/gnm_tensorflow.py index 2ca28350..8f87fe14 100644 --- a/gnm/shape/gnm_tensorflow.py +++ b/gnm/shape/gnm_tensorflow.py @@ -50,6 +50,7 @@ GNMVariant = gnm_specs.GNMVariant GNMBodyPart = gnm_specs.GNMBodyPart GNMLandmarksType = gnm_landmarks.GNMLandmarksType +GNMRemoteSource = gnm_specs.GNMRemoteSource @dataclasses.dataclass(frozen=False, kw_only=True, init=False) diff --git a/gnm/shape/gnm_xnp.py b/gnm/shape/gnm_xnp.py index d7348495..a96108ac 100644 --- a/gnm/shape/gnm_xnp.py +++ b/gnm/shape/gnm_xnp.py @@ -65,7 +65,6 @@ class GNM(gnm_base.GNMBase): Attributes: version: The version of the loaded GNM model. variant: The variant of the loaded GNM model. - cl_number: The CL used to create this model. template_vertex_positions: Vertex positions in the template mesh, (V, 3). template_joint_positions: Joint positions in the template GNM, (J, 3). vertex_identity_basis: The vertex identity basis of the model, (I, V, 3). @@ -187,7 +186,6 @@ def as_original(val): field_converters = { 'version': as_original, 'variant': as_original, - 'cl_number': as_original, 'template_vertex_positions': as_float_array, 'template_joint_positions': as_float_array, 'vertex_identity_basis': as_float_array, diff --git a/gnm/shape/oss_data_loaders/oss_data_loaders.py b/gnm/shape/oss_data_loaders/oss_data_loaders.py new file mode 100644 index 00000000..6545d352 --- /dev/null +++ b/gnm/shape/oss_data_loaders/oss_data_loaders.py @@ -0,0 +1,287 @@ +# 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. + +"""Remote model loaders and download resolution utilities for OSS releases. + +This module provides loaders for downloading and caching GNM model weights from +public repositories (Hugging Face Hub, Kaggle Models, and HTTP/HTTPS CDNs). +These loaders are isolated from Google3 clients and intended for open-source +consumption, see https://github.com/google/GNM. +""" + +# pylint: disable=protected-access + +import functools +import importlib +import os +from typing import Any +import urllib.request + +from absl import logging +from etils import epath +from gnm.shape import gnm_data_loader +from gnm.shape.data.versions import gnm_specs + +DEFAULT_KAGGLE_HANDLE_PREFIX = 'google/gnm-{major}/other' +DEFAULT_HF_REPO = 'google/gnm-{major}' +DEFAULT_HF_CDN_BASE_URL = ( + 'https://huggingface.co/google/gnm-{major}/resolve/main' +) + + +def get_default_gnm_cache_dir() -> epath.Path: + """Returns the default directory for caching downloaded GNM models.""" + if env_cache := os.getenv('GNM_CACHE_DIR'): + return epath.Path(env_cache) + if xdg_cache := os.getenv('XDG_CACHE_HOME'): + return epath.Path(xdg_cache) / 'gnm' / 'models' + return epath.Path(os.path.expanduser('~/.cache/gnm/models')) + + +def _download_file( + url: str, + destination: epath.Path, + timeout: int = 120, +) -> epath.Path: + """Downloads a remote file via HTTP/HTTPS to a destination path atomically.""" + destination.parent.mkdir(parents=True, exist_ok=True) + temp_destination = destination.with_suffix( + f'{destination.suffix}.tmp.{os.getpid()}' + ) + logging.info('Downloading %s to %s...', url, destination) + req = urllib.request.Request( + url, + headers={'User-Agent': 'gnm-client-python'}, + ) + try: + with ( + urllib.request.urlopen(req, timeout=timeout) as response, + temp_destination.open('wb') as out_f, + ): + while True: + chunk = response.read(64 * 1024) + if not chunk: + break + out_f.write(chunk) + except Exception: + if temp_destination.exists(): + temp_destination.unlink() + raise + + temp_destination.replace(destination) + return destination + + +def _resolve_remote_model_file( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + cache_dir: epath.Path, + force_download: bool = False, +) -> epath.Path: + """Downloads the model file via HTTP/HTTPS from the official CDN.""" + version_dir_name = gnm_data_loader._get_version_dir_name(version) + major_tag = version_dir_name.split('_', maxsplit=1)[0] + model_file_name = gnm_data_loader._get_model_filename(version, variant) + cached_file = cache_dir / version_dir_name / model_file_name + if cached_file.exists() and not force_download: + return cached_file + + cdn_base = DEFAULT_HF_CDN_BASE_URL.format(major=major_tag) + cdn_url = f'{cdn_base}/{version_dir_name}/{model_file_name}' + try: + return _download_file(cdn_url, cached_file) + except Exception as e: + raise FileNotFoundError( + f'Could not download GNM model file for version {version} and variant' + f' {variant} from CDN URL: {cdn_url}.\nError: {e}' + ) from e + + +def _resolve_huggingface_model_file( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + cache_dir: epath.Path, + force_download: bool = False, +) -> epath.Path: + """Resolves model file from HF Hub via huggingface_hub SDK or CDN fallback.""" + version_dir_name = gnm_data_loader._get_version_dir_name(version) + major_tag = version_dir_name.split('_', maxsplit=1)[0] + model_file_name = gnm_data_loader._get_model_filename(version, variant) + filename = f'{version_dir_name}/{model_file_name}' + effective_repo_id = f'google/gnm-{major_tag}' + revision = 'main' + + try: + huggingface_hub = importlib.import_module('huggingface_hub') + downloaded_path = huggingface_hub.hf_hub_download( + repo_id=effective_repo_id, + filename=filename, + revision=revision, + cache_dir=str(cache_dir), + force_download=force_download, + ) + return epath.Path(downloaded_path) + except ImportError: + cdn_url = ( + f'https://huggingface.co/{effective_repo_id}/resolve/{revision}/' + f'{filename}' + ) + cached_file = cache_dir / effective_repo_id.replace('/', '_') / filename + if cached_file.exists() and not force_download: + return cached_file + try: + return _download_file(cdn_url, cached_file) + except Exception as e: + raise FileNotFoundError( + f'Could not download GNM model file for version {version} and variant' + f' {variant} from Hugging Face CDN URL: {cdn_url}.\nError: {e}' + ) from e + + +def _resolve_kaggle_model_file( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + cache_dir: epath.Path, + force_download: bool = False, +) -> epath.Path: + """Resolves model file from Kaggle Models using kagglehub SDK.""" + del cache_dir # kagglehub manages its own internal cache directory. + try: + kagglehub = importlib.import_module('kagglehub') + except ImportError as e: + raise ImportError( + 'Loading from Kaggle requires kagglehub. Run: pip install kagglehub' + ) from e + + version_dir_name = gnm_data_loader._get_version_dir_name(version) + major_tag = version_dir_name.split('_', maxsplit=1)[0] + model_file_name = gnm_data_loader._get_model_filename(version, variant) + npz_stem = model_file_name.removesuffix('.npz') + variation_slug = f'{npz_stem}_{version_dir_name}' + kaggle_handle = f'google/gnm-{major_tag}/other/{variation_slug}' + + downloaded_path = kagglehub.model_download( + kaggle_handle, + path=model_file_name, + force_download=force_download, + ) + result_path = epath.Path(downloaded_path) + if result_path.is_dir(): + result_path = result_path / model_file_name + return result_path + + +@functools.lru_cache +def load_model_from_remote( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + source: gnm_specs.GNMRemoteSource = gnm_specs.GNMRemoteSource.HTTP, + *, + cache_dir: epath.PathLike | None = None, + force_download: bool = False, +) -> dict[str, Any]: + """Loads GNM model data via remote repository (HTTP, Hugging Face, or Kaggle). + + Args: + version: GNM major version. + variant: GNM model variant. + source: Remote repository source. Defaults to HTTP (CDN direct download). + cache_dir: Custom local cache directory (Path or str). Defaults to + `~/.cache/gnm/models/`. + force_download: If True, forces redownload even if cached locally. + + Returns: + A dictionary containing the standardized GNM model data. + + Raises: + ImportError: If the required SDK is not installed (for Kaggle). + FileNotFoundError: If the model file cannot be downloaded. + ValueError: If validation of the model data fails or source is invalid. + """ + cache_path = ( + epath.Path(cache_dir) + if cache_dir is not None + else get_default_gnm_cache_dir() + ) + if source in ( + gnm_specs.GNMRemoteSource.HTTP, + 'http', + ): + model_file = _resolve_remote_model_file( + version, variant, cache_path, force_download + ) + elif source in ( + gnm_specs.GNMRemoteSource.HUGGING_FACE, + 'huggingface', + 'hf', + ): + model_file = _resolve_huggingface_model_file( + version, variant, cache_path, force_download + ) + elif source in ( + gnm_specs.GNMRemoteSource.KAGGLE, + 'kaggle', + ): + model_file = _resolve_kaggle_model_file( + version, variant, cache_path, force_download + ) + else: + raise ValueError(f'Unsupported remote source: {source}') + + logging.info( + 'Loading GNM model version %s, variant %s from %s: %s', + version, + variant, + source, + model_file, + ) + return gnm_data_loader._load_model_dict_from_file( + model_file, version, variant + ) + + +@functools.lru_cache +def load_model_from_huggingface( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + *, + cache_dir: epath.PathLike | None = None, + force_download: bool = False, +) -> dict[str, Any]: + """Loads GNM model data from Hugging Face Hub.""" + return load_model_from_remote( + version=version, + variant=variant, + source=gnm_specs.GNMRemoteSource.HUGGING_FACE, + cache_dir=cache_dir, + force_download=force_download, + ) + + +@functools.lru_cache +def load_model_from_kaggle( + version: gnm_specs.GNMMajorVersion, + variant: gnm_specs.GNMVariant, + *, + cache_dir: epath.PathLike | None = None, + force_download: bool = False, +) -> dict[str, Any]: + """Loads GNM model data from Kaggle Models.""" + return load_model_from_remote( + version=version, + variant=variant, + source=gnm_specs.GNMRemoteSource.KAGGLE, + cache_dir=cache_dir, + force_download=force_download, + ) diff --git a/gnm/shape/oss_data_loaders/oss_data_loaders_test.py b/gnm/shape/oss_data_loaders/oss_data_loaders_test.py new file mode 100644 index 00000000..86c847f2 --- /dev/null +++ b/gnm/shape/oss_data_loaders/oss_data_loaders_test.py @@ -0,0 +1,352 @@ +# 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. + +"""Unit tests verifying OSS GNM model data loaders. + +Tests remote file resolution and download logic for Hugging Face and Kaggle. +""" + +# pylint: disable=protected-access + +import email.message +import io +import os +from typing import Any +from unittest import mock +import urllib.error +import urllib.request + +from absl.testing import absltest +from absl.testing import parameterized +from etils import epath +from gnm.shape.data.versions import gnm_specs +from gnm.shape.oss_data_loaders import oss_data_loaders +import numpy as np + + +def _get_dummy_gnm_data_dict() -> dict[str, Any]: + """Returns a dummy GNM data dictionary.""" + return { + 'version': '3.0', + 'variant': 'head', + 'template_vertex_positions': np.zeros((1, 3)), + 'template_joint_positions': np.zeros((1, 3)), + 'vertex_identity_basis': np.zeros((1, 1, 3)), + 'joint_identity_basis': np.zeros((1, 1, 3)), + 'expression_basis': np.zeros((1, 1, 3)), + 'identity_names': ['id1'], + 'joint_names': ['joint1'], + 'expression_names': ['exp1'], + 'joint_parent_indices': np.array([0]), + 'skinning_weights': np.zeros((1, 1)), + 'quads': np.zeros((1, 4)), + 'triangles': np.zeros((1, 3)), + 'quad_uvs': np.zeros((1, 4, 2)), + 'triangle_uvs': np.zeros((1, 3, 2)), + 'mesh_component_names': ['part1'], + 'mirror_indices': np.array([0]), + 'joint_regressor': np.zeros((1, 1)), + 'pose_correctives_regressor': np.zeros((9, 3)), + 'bone_aligned_template_joint_orientations': np.zeros((1, 3, 3)), + 'vertex_groups': np.zeros((1, 1)), + 'vertex_group_names': ['group1'], + } + + +class OSSDataLoadersTest(parameterized.TestCase): + """Tests for remote model loading and caching in oss_data_loaders.""" + + def setUp(self): + super().setUp() + self.temp_dir = epath.Path(self.create_tempdir().full_path) + self.dummy_gnm_data_dict = _get_dummy_gnm_data_dict() + + # Save a dummy npz in temp_dir. + buffer = io.BytesIO() + np.savez_compressed(buffer, **self.dummy_gnm_data_dict) + self.dummy_npz_bytes = buffer.getvalue() + + def test_get_default_gnm_cache_dir(self): + with mock.patch.dict(os.environ, {'GNM_CACHE_DIR': '/custom/gnm/cache'}): + self.assertEqual( + oss_data_loaders.get_default_gnm_cache_dir(), + epath.Path('/custom/gnm/cache'), + ) + + with mock.patch.dict( + os.environ, + {'XDG_CACHE_HOME': '/custom/xdg/cache'}, + clear=True, + ): + self.assertEqual( + oss_data_loaders.get_default_gnm_cache_dir(), + epath.Path('/custom/xdg/cache/gnm/models'), + ) + + def test_download_file_success(self): + dest_path = self.temp_dir / 'downloaded_model.npz' + mock_response = io.BytesIO(self.dummy_npz_bytes) + + with mock.patch.object( + urllib.request, 'urlopen', return_value=mock_response + ): + result_path = oss_data_loaders._download_file( + 'https://huggingface.co/google/gnm-v3/resolve/main/v3_0/gnm_head.npz', + dest_path, + ) + self.assertEqual(result_path, dest_path) + self.assertTrue(dest_path.exists()) + self.assertEqual(dest_path.read_bytes(), self.dummy_npz_bytes) + + def test_download_file_http_error(self): + dest_path = self.temp_dir / 'fail_model.npz' + with mock.patch.object( + urllib.request, + 'urlopen', + side_effect=urllib.error.HTTPError( + 'https://example.com/not_found.npz', + 404, + 'Not Found', + email.message.Message(), + None, + ), + ): + with self.assertRaises(urllib.error.HTTPError): + oss_data_loaders._download_file( + 'https://example.com/not_found.npz', dest_path + ) + self.assertFalse(dest_path.exists()) + + def test_load_model_from_remote_with_version_and_caching(self): + dest_cache_file = self.temp_dir / 'v3_0' / 'gnm_head.npz' + + def _fake_download(url, dest): + del url + dest.parent.mkdir(parents=True, exist_ok=True) + dest.write_bytes(self.dummy_npz_bytes) + return dest + + with mock.patch.object( + oss_data_loaders, '_download_file', side_effect=_fake_download + ) as mock_download: + # First load: triggers download + data1 = oss_data_loaders.load_model_from_remote( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data1, dict) + self.assertEqual(mock_download.call_count, 1) + self.assertTrue(dest_cache_file.exists()) + + # Second load: uses cached file directly, does not re-download + data2 = oss_data_loaders.load_model_from_remote( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data2, dict) + self.assertEqual(mock_download.call_count, 1) + + # Force download: re-downloads even if cached + data3 = oss_data_loaders.load_model_from_remote( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + force_download=True, + ) + self.assertIsInstance(data3, dict) + self.assertEqual(mock_download.call_count, 2) + + def test_load_model_from_remote_with_str_cache_dir(self): + dest_cache_file = self.temp_dir / 'v3_0' / 'gnm_head.npz' + + def _fake_download(url, dest): + del url + dest.parent.mkdir(parents=True, exist_ok=True) + dest.write_bytes(self.dummy_npz_bytes) + return dest + + with mock.patch.object( + oss_data_loaders, '_download_file', side_effect=_fake_download + ): + data = oss_data_loaders.load_model_from_remote( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=str(self.temp_dir), + ) + self.assertIsInstance(data, dict) + self.assertTrue(dest_cache_file.exists()) + + def test_load_model_from_huggingface(self): + dest_file = self.temp_dir / 'v3_0' / 'gnm_head.npz' + dest_file.parent.mkdir(parents=True, exist_ok=True) + dest_file.write_bytes(self.dummy_npz_bytes) + + with mock.patch.object( + oss_data_loaders, + '_resolve_huggingface_model_file', + return_value=dest_file, + ) as mock_resolve: + data = oss_data_loaders.load_model_from_huggingface( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data, dict) + mock_resolve.assert_called_once_with( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + self.temp_dir, + False, + ) + + def test_load_model_from_remote_huggingface(self): + dest_file = self.temp_dir / 'v3_0' / 'gnm_head.npz' + dest_file.parent.mkdir(parents=True, exist_ok=True) + dest_file.write_bytes(self.dummy_npz_bytes) + + with mock.patch.object( + oss_data_loaders, + '_resolve_huggingface_model_file', + return_value=dest_file, + ) as mock_resolve: + data = oss_data_loaders.load_model_from_remote( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + source=gnm_specs.GNMRemoteSource.HUGGING_FACE, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data, dict) + mock_resolve.assert_called_once_with( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + self.temp_dir, + False, + ) + + def test_load_model_from_kaggle(self): + dest_file = self.temp_dir / 'v3_0' / 'gnm_head.npz' + dest_file.parent.mkdir(parents=True, exist_ok=True) + dest_file.write_bytes(self.dummy_npz_bytes) + + with mock.patch.object( + oss_data_loaders, + '_resolve_kaggle_model_file', + return_value=dest_file, + ) as mock_resolve: + data = oss_data_loaders.load_model_from_kaggle( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data, dict) + mock_resolve.assert_called_once_with( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + self.temp_dir, + False, + ) + + def test_load_model_from_remote_kaggle(self): + dest_file = self.temp_dir / 'v3_0' / 'gnm_head.npz' + dest_file.parent.mkdir(parents=True, exist_ok=True) + dest_file.write_bytes(self.dummy_npz_bytes) + + with mock.patch.object( + oss_data_loaders, + '_resolve_kaggle_model_file', + return_value=dest_file, + ) as mock_resolve: + data = oss_data_loaders.load_model_from_remote( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + source=gnm_specs.GNMRemoteSource.KAGGLE, + cache_dir=self.temp_dir, + ) + self.assertIsInstance(data, dict) + mock_resolve.assert_called_once_with( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + self.temp_dir, + False, + ) + + def test_load_model_from_remote_invalid_source(self): + with self.assertRaisesRegex(ValueError, 'Unsupported remote source'): + oss_data_loaders.load_model_from_remote( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + source='invalid_source', # pyrefly: ignore[bad-argument-type] + cache_dir=self.temp_dir, + ) + + def test_resolve_huggingface_model_file_sdk(self): + mock_hf = mock.MagicMock() + mock_hf.hf_hub_download.return_value = '/downloaded/path/gnm_head.npz' + with mock.patch.object( + oss_data_loaders.importlib, 'import_module', return_value=mock_hf + ): + res = oss_data_loaders._resolve_huggingface_model_file( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + ) + self.assertEqual(res, epath.Path('/downloaded/path/gnm_head.npz')) + mock_hf.hf_hub_download.assert_called_once_with( + repo_id='google/gnm-v3', + filename='v3_0/gnm_head.npz', + revision='main', + cache_dir=str(self.temp_dir), + force_download=False, + ) + + def test_resolve_huggingface_model_file_fallback_error(self): + with mock.patch.object( + oss_data_loaders.importlib, + 'import_module', + side_effect=ImportError('No HF'), + ): + with mock.patch.object( + oss_data_loaders, + '_download_file', + side_effect=RuntimeError('Network unreachable'), + ): + with self.assertRaises(FileNotFoundError): + oss_data_loaders._resolve_huggingface_model_file( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + ) + + def test_resolve_kaggle_model_file(self): + fake_kagglehub = mock.MagicMock() + expected_file = self.temp_dir / 'gnm_head.npz' + fake_kagglehub.model_download.return_value = str(expected_file) + with mock.patch.dict('sys.modules', {'kagglehub': fake_kagglehub}): + path = oss_data_loaders._resolve_kaggle_model_file( + gnm_specs.GNMMajorVersion.V3, + gnm_specs.GNMVariant.HEAD, + cache_dir=self.temp_dir, + ) + fake_kagglehub.model_download.assert_called_once_with( + 'google/gnm-v3/other/gnm_head_v3_0', + path='gnm_head.npz', + force_download=False, + ) + self.assertEqual(path, expected_file) + + +if __name__ == '__main__': + absltest.main()