Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions gnm/shape/data/versions/gnm_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'
Binary file modified gnm/shape/data/versions/v3_0/gnm_head.npz
Binary file not shown.
32 changes: 29 additions & 3 deletions gnm/shape/demos/gnm_head_demo.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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,
Expand Down
20 changes: 18 additions & 2 deletions gnm/shape/demos/semantic_gnm_demo.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -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"
],
Expand Down
33 changes: 33 additions & 0 deletions gnm/shape/gnm_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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."""
Expand Down
75 changes: 71 additions & 4 deletions gnm/shape/gnm_base_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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()
72 changes: 56 additions & 16 deletions gnm/shape/gnm_data_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@

"""GNM data loader."""

# from collections.abc import Mapping, Sequence
from collections.abc import Sequence
import functools
from typing import Any
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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]]:
Expand All @@ -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


Expand Down
Loading
Loading