Skip to content
Merged
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 checkpoint/CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
209 changes: 209 additions & 0 deletions checkpoint/orbax/checkpoint/_src/path/fs_probe.py
Original file line number Diff line number Diff line change
@@ -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)))
175 changes: 175 additions & 0 deletions checkpoint/orbax/checkpoint/_src/path/fs_probe_test.py
Original file line number Diff line number Diff line change
@@ -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()
Loading
Loading