From 2d2225be9ae5143112a699abfbca2a8318dd23a7 Mon Sep 17 00:00:00 2001 From: Daniel Ng Date: Mon, 21 Sep 2026 15:29:36 -0700 Subject: [PATCH] Internal Change. PiperOrigin-RevId: 985549374 --- checkpoint/CHANGELOG.md | 5 + .../orbax/checkpoint/_src/path/fs_probe.py | 209 ++++++++++++++++++ .../checkpoint/_src/path/fs_probe_test.py | 175 +++++++++++++++ .../experimental/v1/_src/context/options.py | 3 + .../v1/_src/layout/orbax_layout.py | 10 + .../experimental/v1/_src/layout/registry.py | 186 +++++++++++----- .../v1/_src/layout/registry_test.py | 103 +++++++++ .../v1/_src/layout/safetensors_layout.py | 11 + 8 files changed, 652 insertions(+), 50 deletions(-) create mode 100644 checkpoint/orbax/checkpoint/_src/path/fs_probe.py create mode 100644 checkpoint/orbax/checkpoint/_src/path/fs_probe_test.py diff --git a/checkpoint/CHANGELOG.md b/checkpoint/CHANGELOG.md index 24e1cfd736..75f4cac74e 100644 --- a/checkpoint/CHANGELOG.md +++ b/checkpoint/CHANGELOG.md @@ -7,6 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- #v1 Add `CheckpointLayout.AUTO` to detect the checkpoint layout from + filesystem markers when loading. + ## [0.12.5] - 2026-09-17 ### Added diff --git a/checkpoint/orbax/checkpoint/_src/path/fs_probe.py b/checkpoint/orbax/checkpoint/_src/path/fs_probe.py new file mode 100644 index 0000000000..95dd2d5399 --- /dev/null +++ b/checkpoint/orbax/checkpoint/_src/path/fs_probe.py @@ -0,0 +1,209 @@ +# Copyright 2026 The Orbax Authors. +# +# 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 +# +# http://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. + +"""Batched filesystem probing utilities for checkpoint inspection. + +A shared path utility library providing low-latency directory indexing and +concurrent probe primitives, especially used by checkpoint layout and format +detection routines. + +Probing individual marker files sequentially on distributed or cloud +filesystems incurs significant RPC round-trip latency. This module minimizes +filesystem round trips through two core primitives: + +1. **Ask once, answer many:** A single non-recursive directory scan answers all + marker existence checks directly under a directory with one round trip. + `DirectoryIndex` snapshots that listing for fast membership lookups. +2. **Ask in parallel:** When multiple subdirectories or paths must be checked + concurrently, `exists_many`, `is_dir_many`, and `index_directories` issue + checks concurrently via asyncio rather than serially. + +Note: Directory listings are strictly non-recursive to avoid traversing large +tensor or chunk subtrees. +""" + +from __future__ import annotations + +import asyncio + +from etils import epath +from orbax.checkpoint._src.path import async_path + + +class DirectoryIndex: + """The immediate contents of one directory, fetched in a single round trip.""" + + def __init__( + self, + path: epath.Path, + exists: bool = False, + is_directory: bool = False, + names: frozenset[str] = frozenset(), + ): + """Initializes the directory index.""" + self._path = path + self._exists = exists + self._is_directory = is_directory + self._names = names + + def path(self) -> epath.Path: + """Returns the directory path that was indexed.""" + return self._path + + def exists(self) -> bool: + """Returns whether the path exists.""" + return self._exists + + def is_directory(self) -> bool: + """Returns whether the path is a directory.""" + return self._is_directory + + def listable(self) -> bool: + """Returns whether the path exists and is a directory.""" + return self._exists and self._is_directory + + def names(self) -> frozenset[str]: + """Returns immediate child names.""" + return self._names + + def has(self, name: str) -> bool: + """Returns whether a child with exactly this name exists.""" + return name in self._names + + def has_any(self, *candidates: str) -> bool: + """Returns whether any of the named children exist.""" + return any(candidate in self._names for candidate in candidates) + + def matching(self, *prefixes: str) -> list[str]: + """Returns sorted child names starting with any of the given prefixes.""" + return sorted(name for name in self._names if name.startswith(prefixes)) + + def with_suffix(self, suffix: str) -> list[str]: + """Returns sorted child names ending with the given suffix.""" + return sorted(name for name in self._names if name.endswith(suffix)) + + def present(self, candidates: tuple[str, ...]) -> list[str]: + """Returns the candidates that exist, preserving candidate order.""" + return [candidate for candidate in candidates if candidate in self._names] + + def __repr__(self) -> str: + return ( + f"DirectoryIndex(path={self._path!r}, exists={self._exists!r}, " + f"is_directory={self._is_directory!r}, names={self._names!r})" + ) + + def __eq__(self, other: object) -> bool: + if not isinstance(other, DirectoryIndex): + return False + return ( + self._path == other._path + and self._exists == other._exists + and self._is_directory == other._is_directory + and self._names == other._names + ) + + +_MISSING = frozenset() + + +async def index_directory(path: epath.Path) -> DirectoryIndex: + """Lists one directory asynchronously, discovering child names.""" + try: + entries = await async_path.iterdir(path) + names = frozenset(entry.name for entry in entries) + return DirectoryIndex( + path=path, exists=True, is_directory=True, names=names + ) + except FileNotFoundError: + return DirectoryIndex( + path=path, exists=False, is_directory=False, names=_MISSING + ) + except NotADirectoryError: + return DirectoryIndex( + path=path, exists=True, is_directory=False, names=_MISSING + ) + except Exception: # pylint: disable=broad-exception-caught + # Probing remote/cloud filesystems may raise driver-specific exceptions. + try: + exists = await async_path.exists(path) + is_dir = await async_path.is_dir(path) if exists else False + return DirectoryIndex( + path=path, exists=exists, is_directory=is_dir, names=_MISSING + ) + except Exception: # pylint: disable=broad-exception-caught + # Fall back to missing index if exists/is_dir probe fails. + return DirectoryIndex( + path=path, exists=False, is_directory=False, names=_MISSING + ) + + +async def index_directories( + paths: tuple[epath.Path, ...], +) -> tuple[DirectoryIndex, ...]: + """Lists several directories concurrently. + + Args: + paths: Directories to list. + + Returns: + Indexes positionally aligned with `paths`. + """ + return tuple(await asyncio.gather(*(index_directory(p) for p in paths))) + + +async def exists_many( + paths: tuple[epath.Path, ...], +) -> tuple[bool, ...]: + """Checks several paths for existence concurrently. + + Prefer answering from a `DirectoryIndex` when the paths share a parent; use + this only when they do not. + + Args: + paths: Paths to check. + + Returns: + Booleans positionally aligned with `paths`. + """ + + async def _safe_exists(target_path: epath.Path) -> bool: + try: + return await async_path.exists(target_path) + except Exception: # pylint: disable=broad-exception-caught + # Remote/cloud filesystem driver exceptions treat path as nonexistent. + return False + + return tuple(await asyncio.gather(*(_safe_exists(p) for p in paths))) + + +async def is_dir_many( + paths: tuple[epath.Path, ...], +) -> tuple[bool, ...]: + """Checks several paths for directory-ness concurrently. + + Args: + paths: Paths to check. + + Returns: + Booleans positionally aligned with `paths`. + """ + + async def _safe_is_dir(target_path: epath.Path) -> bool: + try: + return await async_path.is_dir(target_path) + except Exception: # pylint: disable=broad-exception-caught + # Remote/cloud filesystem driver exceptions treat path as non-directory. + return False + + return tuple(await asyncio.gather(*(_safe_is_dir(p) for p in paths))) diff --git a/checkpoint/orbax/checkpoint/_src/path/fs_probe_test.py b/checkpoint/orbax/checkpoint/_src/path/fs_probe_test.py new file mode 100644 index 0000000000..5bb59a3b29 --- /dev/null +++ b/checkpoint/orbax/checkpoint/_src/path/fs_probe_test.py @@ -0,0 +1,175 @@ +# Copyright 2026 The Orbax Authors. +# +# 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 +# +# http://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 for fs_probe.""" + +from __future__ import annotations + +import unittest + +from absl.testing import absltest +from absl.testing import parameterized +from etils import epath +from orbax.checkpoint._src.path import fs_probe + + +class DirectoryIndexTest( + parameterized.TestCase, unittest.IsolatedAsyncioTestCase +): + + def setUp(self): + super().setUp() + self.root = epath.Path(self.create_tempdir().full_path) + + def _touch(self, *names: str) -> None: + for name in names: + (self.root / name).write_text("") + + async def test_lists_immediate_children_only(self): + self._touch("a.txt", "b.txt") + nested = self.root / "sub" + nested.mkdir() + (nested / "hidden.txt").write_text("") + + index = await fs_probe.index_directory(self.root) + + self.assertEqual(index.path(), self.root) + self.assertTrue(index.exists()) + self.assertTrue(index.is_directory()) + self.assertTrue(index.listable()) + self.assertEqual(index.names(), frozenset({"a.txt", "b.txt", "sub"})) + self.assertNotIn("hidden.txt", index.names()) + + async def test_missing_directory(self): + index = await fs_probe.index_directory(self.root / "nope") + + self.assertFalse(index.exists()) + self.assertFalse(index.is_directory()) + self.assertFalse(index.listable()) + self.assertEmpty(index.names()) + self.assertFalse(index.has("anything")) + + async def test_file_as_directory(self): + file_path = self.root / "a_file.txt" + file_path.write_text("content") + index = await fs_probe.index_directory(file_path) + + self.assertTrue(index.exists()) + self.assertFalse(index.is_directory()) + self.assertFalse(index.listable()) + self.assertEmpty(index.names()) + + async def test_has_and_has_any(self): + self._touch("manifest.ocdbt", "_METADATA") + index = await fs_probe.index_directory(self.root) + + self.assertTrue(index.has("manifest.ocdbt")) + self.assertFalse(index.has("manifest")) + self.assertTrue(index.has_any("missing", "_METADATA")) + self.assertFalse(index.has_any("missing", "also_missing")) + + async def test_matching_returns_sorted_prefix_hits(self): + self._touch( + "ocdbt.process_1", + "ocdbt.process_0", + "manifest.ocdbt", + ) + index = await fs_probe.index_directory(self.root) + + self.assertEqual( + index.matching("ocdbt.process_"), + ["ocdbt.process_0", "ocdbt.process_1"], + ) + self.assertEmpty(index.matching("no_such_prefix")) + + async def test_matching_accepts_several_prefixes(self): + self._touch("a_one", "b_two", "c_three") + index = await fs_probe.index_directory(self.root) + + self.assertEqual(index.matching("a_", "c_"), ["a_one", "c_three"]) + + async def test_with_suffix(self): + self._touch("model.safetensors", "other.safetensors", "notes.txt") + index = await fs_probe.index_directory(self.root) + + self.assertEqual( + index.with_suffix(".safetensors"), + ["model.safetensors", "other.safetensors"], + ) + + async def test_present_preserves_candidate_order(self): + self._touch("second", "first") + index = await fs_probe.index_directory(self.root) + + self.assertEqual( + index.present(("first", "absent", "second")), ["first", "second"] + ) + + +class BatchProbeTest(parameterized.TestCase, unittest.IsolatedAsyncioTestCase): + + def setUp(self): + super().setUp() + self.root = epath.Path(self.create_tempdir().full_path) + + async def test_index_directories_is_positionally_aligned(self): + (self.root / "a").mkdir() + (self.root / "a" / "x").write_text("") + (self.root / "b").mkdir() + + indexes = await fs_probe.index_directories( + (self.root / "a", self.root / "missing", self.root / "b") + ) + + self.assertLen(indexes, 3) + self.assertEqual(indexes[0].names(), frozenset({"x"})) + self.assertTrue(indexes[0].exists()) + self.assertTrue(indexes[0].is_directory()) + + self.assertFalse(indexes[1].exists()) + self.assertFalse(indexes[1].is_directory()) + self.assertFalse(indexes[1].listable()) + + self.assertTrue(indexes[2].exists()) + self.assertTrue(indexes[2].is_directory()) + self.assertTrue(indexes[2].listable()) + self.assertEmpty(indexes[2].names()) + + async def test_exists_many(self): + (self.root / "here").write_text("") + + self.assertEqual( + await fs_probe.exists_many((self.root / "here", self.root / "gone")), + (True, False), + ) + + async def test_is_dir_many(self): + (self.root / "dir").mkdir() + (self.root / "file").write_text("") + + self.assertEqual( + await fs_probe.is_dir_many( + (self.root / "dir", self.root / "file", self.root / "gone") + ), + (True, False, False), + ) + + async def test_empty_input_does_no_work(self): + self.assertEmpty(await fs_probe.index_directories(())) + self.assertEmpty(await fs_probe.exists_many(())) + self.assertEmpty(await fs_probe.is_dir_many(())) + + +if __name__ == "__main__": + absltest.main() diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/context/options.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/context/options.py index b1e7fe8d17..2349316fca 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/context/options.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/context/options.py @@ -686,9 +686,12 @@ class CheckpointLayout(enum.Enum): Currently supported layouts are: + AUTO: Automatically detects layout from filesystem markers when + loading. Resolves to ORBAX on save. Opt-in; default remains ORBAX. ORBAX: Orbax's own layout. SAFETENSORS: https://huggingface.co/docs/safetensors/en/index """ ORBAX = 'orbax' SAFETENSORS = 'safetensors' + AUTO = 'auto' diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/orbax_layout.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/orbax_layout.py index ef0e0072d2..96c17d94f4 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/orbax_layout.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/orbax_layout.py @@ -23,6 +23,7 @@ from orbax.checkpoint._src.metadata import step_metadata_serialization from orbax.checkpoint._src.multihost import multihost from orbax.checkpoint._src.path import async_path +from orbax.checkpoint._src.path import fs_probe from orbax.checkpoint._src.path import temporary_paths from orbax.checkpoint.experimental.v1._src.context import context as context_lib from orbax.checkpoint.experimental.v1._src.handlers import registration @@ -60,6 +61,15 @@ class CheckpointVersion(enum.Enum): _ZARRAY_FILE = ".zarray" +async def matches_markers(index: fs_probe.DirectoryIndex) -> bool: + """Returns whether `index` shows the discriminating markers of this layout.""" + return index.has_any( + ORBAX_CHECKPOINT_INDICATOR_FILE, + CHECKPOINT_METADATA, + PYTREE_METADATA_FILE, + ) + + async def checkpoint_version(path: path_types.PathLike) -> CheckpointVersion: """Returns the checkpoint version of the given path.""" if await has_indicator_file(path): # pyrefly: ignore[bad-argument-type] diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry.py index be9937670a..5ab815b524 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry.py @@ -16,9 +16,11 @@ from __future__ import annotations import asyncio +from typing import Any from absl import logging from orbax.checkpoint._src import asyncio_utils +from orbax.checkpoint._src.path import fs_probe from orbax.checkpoint.experimental.v1._src.context import context as context_lib from orbax.checkpoint.experimental.v1._src.context import options as options_lib from orbax.checkpoint.experimental.v1._src.layout import checkpoint_layout @@ -39,6 +41,14 @@ async def _is_orbax_checkpoint_async(path: path_types.PathLike) -> bool: + """Checks asynchronously whether the path is an Orbax checkpoint. + + Args: + path: Path to the checkpoint to check. + + Returns: + True if the checkpoint matches any registered Orbax layout class. + """ ctx = context_lib.get_context() path = ctx.file_options.path_class(path) @@ -51,15 +61,87 @@ async def _is_orbax_checkpoint_async(path: path_types.PathLike) -> bool: def is_orbax_checkpoint(path: path_types.PathLike) -> bool: - """Returns True if the path is an Orbax checkpoint.""" + """Returns True if the path is an Orbax checkpoint. + + Args: + path: Path to the checkpoint to check. + + Returns: + True if the path is recognized as an Orbax checkpoint. + """ return asyncio_utils.run_sync(_is_orbax_checkpoint_async(path)) +async def detect_layout( + path: path_types.PathLike, + *, + root_index: fs_probe.DirectoryIndex | None = None, +) -> CheckpointLayoutEnum: + """Detects the checkpoint layout from filesystem markers. + + Layouts are evaluated concurrently for primary markers. Orbax is checked first + or prioritized according to specific layout markers, falling back to + Safetensors if primary layouts do not match. + + Args: + path: The path to the checkpoint directory or file. + root_index: Optional pre-computed directory index of `path`. + + Returns: + The detected CheckpointLayout enum. + + Raises: + InvalidLayoutError: If the path does not match any registered layout. + """ + ctx = context_lib.get_context() + resolved_path = ctx.file_options.path_class(path) + + if root_index is None: + root_index = await fs_probe.index_directory(resolved_path) + + if await orbax_layout.matches_markers(root_index): + return CheckpointLayoutEnum.ORBAX + + + if await safetensors_layout.matches_markers(root_index): + return CheckpointLayoutEnum.SAFETENSORS + + tried = [ + CheckpointLayoutEnum.ORBAX.value, + CheckpointLayoutEnum.SAFETENSORS.value, + ] + raise InvalidLayoutError( + f"Could not auto-detect checkpoint layout at {path}. " + f"Tried layouts: {tried}." + ) + + + + async def get_layout_class( layout_enum: CheckpointLayoutEnum, path: path_types.PathLike | None = None ) -> type[CheckpointLayout]: - """Returns the layout class for the given layout enum.""" + """Returns the layout class for the given layout enum. + + Args: + layout_enum: Layout enum identifying the checkpoint format. + path: Optional checkpoint path used for version or format detection. + + Returns: + The concrete CheckpointLayout class. + + Raises: + ValueError: If layout_enum is not recognized. + ImportError: If the requested layout engine (e.g. Roc) is not linked. + """ match layout_enum: + case CheckpointLayoutEnum.AUTO: + if path is None: + # When saving, there is no existing checkpoint to detect; default to + # ORBAX ("detect on read, always write Orbax"). + return orbax_layout.OrbaxLayout + detected_enum = await detect_layout(path) + return await get_layout_class(detected_enum, path) case CheckpointLayoutEnum.ORBAX: if path is None or ( await orbax_layout.checkpoint_version(path) @@ -70,21 +152,6 @@ async def get_layout_class( return orbax_v0_layout.OrbaxV0Layout case CheckpointLayoutEnum.SAFETENSORS: return safetensors_layout.SafetensorsLayout - case CheckpointLayoutEnum.ROC: - try: - # pylint: disable=g-import-not-at-top - # pytype: disable=import-error - from orbax.checkpoint.experimental.v1._src.layout import roc_layout - # pytype: enable=import-error - # pylint: enable=g-import-not-at-top - except ImportError as e: - raise ImportError( - "Failed to import `roc_layout`; Roc support may not be linked. " - "Please depend on " - "//orbax/checkpoint/experimental/v1:roc_support " - "in your build rule." - ) from e - return roc_layout.RocLayout case _: raise ValueError(f"Unsupported checkpoint layout: {layout_enum}") @@ -125,6 +192,56 @@ async def get_checkpoint_layout( ) from e +async def _resolve_auto_pytree_name( + layout: CheckpointLayout, path: path_types.Path +) -> str | None: + """Discovers and validates a pytree checkpointable name at path. + + Args: + layout: Validated CheckpointLayout instance. + path: Path to the checkpoint directory. + + Returns: + The resolved checkpointable name (str or None). + + Raises: + InvalidLayoutError: If no valid PyTree checkpointable can be found. + """ + names = await layout.get_checkpointable_names(path) + for name in names: + try: + await layout.validate(path, name) + logging.info( + "AUTO resolution mode successfully identified a pytree with" + " checkpointable name '%s' at path '%s'. Attempting to load with" + " this name. If this is not the desired checkpointable, please" + " specify the name explicitly.", + name, + path, + ) + return name + except InvalidLayoutError: + continue + + if not isinstance(layout, orbax_layout.OrbaxLayout): + try: + await layout.validate(path, None) + logging.info( + "AUTO resolution mode successfully identified a pytree at path" + " '%s'. Attempting to load as a flat layout checkpoint with" + " checkpointable_name=None.", + path, + ) + return None + except InvalidLayoutError: + pass + + raise InvalidLayoutError( + "Failed to load checkpoint using AUTO resolution mode on" + f" path='{path}'. No valid PyTree checkpointable found." + ) + + class CheckpointLayoutResolver: """Resolves the layout and pytree name for a checkpoint.""" @@ -169,39 +286,8 @@ async def resolve( layout = await get_checkpoint_layout(path, layout_enum) if pytree_name == checkpoint_layout.AUTO_CHECKPOINTABLE_KEY: - names = await layout.get_checkpointable_names(path) - for name in names: - try: - await layout.validate(path, name) - logging.info( - "AUTO resolution mode successfully identified a pytree with" - " checkpointable name '%s' at path '%s'. Attempting to load with" - " this name. If this is not the desired checkpointable, please" - " specify the name explicitly.", - name, - path, - ) - return cls(path, layout_enum, layout, name) - except InvalidLayoutError: - continue - - if not isinstance(layout, orbax_layout.OrbaxLayout): - try: - await layout.validate(path, None) - logging.info( - "AUTO resolution mode successfully identified a pytree at path" - " '%s'. Attempting to load as a flat layout checkpoint with" - " checkpointable_name=None.", - path, - ) - return cls(path, layout_enum, layout, None) - except InvalidLayoutError: - pass - - raise InvalidLayoutError( - "Failed to load checkpoint using AUTO resolution mode on" - f" path='{path}'. No valid PyTree checkpointable found." - ) from None + resolved_name = await _resolve_auto_pytree_name(layout, path) + return cls(path, layout_enum, layout, resolved_name) await layout.validate(path, pytree_name) return cls(path, layout_enum, layout, pytree_name) diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry_test.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry_test.py index a9c0d3551e..5fa763e5dd 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry_test.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/registry_test.py @@ -12,12 +12,14 @@ # See the License for the specific language governing permissions and # limitations under the License. +import sys import unittest from unittest import mock from absl.testing import absltest from absl.testing import parameterized from etils import epath +import numpy as np from orbax.checkpoint import args from orbax.checkpoint._src.checkpointers import checkpointer from orbax.checkpoint._src.handlers import composite_checkpoint_handler @@ -27,7 +29,9 @@ from orbax.checkpoint.experimental.v1._src.layout import orbax_layout from orbax.checkpoint.experimental.v1._src.layout import orbax_v0_layout from orbax.checkpoint.experimental.v1._src.layout import registry +from orbax.checkpoint.experimental.v1._src.layout import safetensors_layout from orbax.checkpoint.experimental.v1._src.saving import saving +import safetensors.numpy STATE_CHECKPOINTABLE_KEY = checkpoint_layout.STATE_CHECKPOINTABLE_KEY @@ -193,6 +197,105 @@ async def test_lower_level_error_propagation(self): +class DetectLayoutTest( + parameterized.TestCase, unittest.IsolatedAsyncioTestCase +): + + def setUp(self): + super().setUp() + self.root_directory = epath.Path(self.create_tempdir()) + self.v1_directory = self.root_directory / 'v1' + saving.save( + self.v1_directory, + {'a': 1, 'b': 2}, # pyrefly: ignore[bad-argument-type] + ) + self.v0_directory = self.root_directory / 'v0' + ckptr = checkpointer.Checkpointer( + composite_checkpoint_handler.CompositeCheckpointHandler() + ) + ckptr.save( + self.v0_directory, + composite_checkpoint_handler.CompositeArgs( + state=standard_checkpoint_handler.StandardSaveArgs({'a': 1, 'b': 2}) + ), + ) + self.v0_flat_directory = self.root_directory / 'v0_flat' + ckptr_flat = checkpointer.Checkpointer( + standard_checkpoint_handler.StandardCheckpointHandler() + ) + ckptr_flat.save( + self.v0_flat_directory, + standard_checkpoint_handler.StandardSaveArgs({'a': 1, 'b': 2}), + ) + self.safetensors_dir = self.root_directory / 'safetensors_dir' + self.safetensors_dir.mkdir() + self.safetensors_file = self.safetensors_dir / 'model.safetensors' + safetensors.numpy.save_file( + {'weight': np.ones((2, 2))}, str(self.safetensors_file) + ) + + async def test_detect_safetensors_file(self): + detected = await registry.detect_layout(self.safetensors_file) + self.assertEqual(detected, CheckpointLayoutEnum.SAFETENSORS) + + async def test_detect_safetensors_directory(self): + detected = await registry.detect_layout(self.safetensors_dir) + self.assertEqual(detected, CheckpointLayoutEnum.SAFETENSORS) + + async def test_detect_orbax_v1(self): + detected = await registry.detect_layout(self.v1_directory) + self.assertEqual(detected, CheckpointLayoutEnum.ORBAX) + + async def test_detect_orbax_v0(self): + detected = await registry.detect_layout(self.v0_directory) + self.assertEqual(detected, CheckpointLayoutEnum.ORBAX) + + async def test_detect_orbax_v0_flat(self): + detected = await registry.detect_layout(self.v0_flat_directory) + self.assertEqual(detected, CheckpointLayoutEnum.ORBAX) + + async def test_detect_unrecognized_directory_raises(self): + empty_dir = self.root_directory / 'empty' + empty_dir.mkdir() + with self.assertRaises(registry.InvalidLayoutError): + await registry.detect_layout(empty_dir) + + async def test_auto_save_path_resolves_to_orbax(self): + cls = await registry.get_layout_class(CheckpointLayoutEnum.AUTO, path=None) + self.assertEqual(cls, orbax_layout.OrbaxLayout) + + async def test_auto_with_path(self): + cls_v1 = await registry.get_layout_class( + CheckpointLayoutEnum.AUTO, path=self.v1_directory + ) + self.assertEqual(cls_v1, orbax_layout.OrbaxLayout) + + cls_v0 = await registry.get_layout_class( + CheckpointLayoutEnum.AUTO, path=self.v0_directory + ) + self.assertEqual(cls_v0, orbax_v0_layout.OrbaxV0Layout) + + cls_st = await registry.get_layout_class( + CheckpointLayoutEnum.AUTO, path=self.safetensors_dir + ) + self.assertEqual(cls_st, safetensors_layout.SafetensorsLayout) + + async def test_resolver_with_auto(self): + resolver = await registry.CheckpointLayoutResolver.resolve( + self.v1_directory, + CheckpointLayoutEnum.AUTO, + pytree_name=checkpoint_layout.AUTO_CHECKPOINTABLE_KEY, + ) + self.assertEqual(resolver.pytree_name, STATE_CHECKPOINTABLE_KEY) + self.assertIsInstance(resolver.layout, orbax_layout.OrbaxLayout) + + async def test_nonexistent_safetensors_raises_invalid_layout(self): + missing_st = self.root_directory / 'nonexistent.safetensors' + with self.assertRaises(registry.InvalidLayoutError): + await registry.detect_layout(missing_st) + + + class IsOrbaxCheckpointTest(parameterized.TestCase): def setUp(self): diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/safetensors_layout.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/safetensors_layout.py index 8de433a672..1458e00ca1 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/safetensors_layout.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/layout/safetensors_layout.py @@ -60,6 +60,7 @@ from orbax.checkpoint._src.arrays import types as arrays_types from orbax.checkpoint._src.multihost import multihost from orbax.checkpoint._src.path import async_path +from orbax.checkpoint._src.path import fs_probe from orbax.checkpoint._src.serialization import limits from orbax.checkpoint._src.tree import utils as tree_utils from orbax.checkpoint.experimental.v1._src.context import context as context_lib @@ -76,6 +77,16 @@ HEADER_NUM_BYTES = 8 SAFETENSORS_SUFFIX = ".safetensors" + +async def matches_markers(index: fs_probe.DirectoryIndex) -> bool: + """Returns whether `index` shows discriminating markers of this layout.""" + if not index.exists(): + return False + return index.path().suffix == SAFETENSORS_SUFFIX or bool( + index.with_suffix(SAFETENSORS_SUFFIX) + ) + + # Read-planning defaults. `SafetensorsOptions` overrides the over-read ratio # (`max_over_read_ratio`) and the chunk size (`read_chunk_bytes`); the # in-flight budget comes from `MemoryOptions.read_concurrent_bytes`.