diff --git a/checkpoint/CHANGELOG.md b/checkpoint/CHANGELOG.md index 5eb6660a5..a62326b64 100644 --- a/checkpoint/CHANGELOG.md +++ b/checkpoint/CHANGELOG.md @@ -32,6 +32,7 @@ heavy over-read. - Support for `gcs_grpc` driver. - #v1 Add stringent validation for state abstract and concrete leaf types. - Update MTC w/ Pathways to Support Scale Elasticity. +- #v1 Add support for saving leaf values directly. ## [0.12.0] - 2026-06-02 diff --git a/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py b/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py index 8463b33cd..b9d0e05ad 100644 --- a/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py +++ b/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py @@ -641,10 +641,9 @@ async def async_save( """ start_time = time.time() item = args.item - # Reject only zero-leaf items (empty containers, None). A single falsy leaf - # (0, '', False, zero-array) is a valid one-leaf tree and must be allowed. - if not jax.tree.leaves(item) and not item: - raise ValueError('Found empty item.') + # `None` is the only unsaveable item: it carries no structure to recover. + if item is None: + raise ValueError('None is not saveable.') save_args = args.save_args ocdbt_target_data_file_size = args.ocdbt_target_data_file_size custom_metadata = args.custom_metadata diff --git a/checkpoint/orbax/checkpoint/_src/handlers/pytree_checkpoint_handler_test.py b/checkpoint/orbax/checkpoint/_src/handlers/pytree_checkpoint_handler_test.py index 17dc3c8d6..0e0ee9b76 100644 --- a/checkpoint/orbax/checkpoint/_src/handlers/pytree_checkpoint_handler_test.py +++ b/checkpoint/orbax/checkpoint/_src/handlers/pytree_checkpoint_handler_test.py @@ -987,6 +987,7 @@ def make_params(): [], [1, [], 2], {'a': [], 'b': 3}, + (), ), save_args=( None, @@ -1011,8 +1012,11 @@ def test_empty_data( with self.ocdbt_checkpoint_handler( use_ocdbt, array_metadata_store=array_metadata_store ) as checkpoint_handler: - if not data: - with self.assertRaisesRegex(ValueError, 'Found empty item'): + # Only `None` is unsaveable. Empty containers ({}, [], ()) are stored as + # a metadata-only entry at the root keypath and restore to the same + # empty container. + if data is None: + with self.assertRaisesRegex(ValueError, 'None is not saveable'): checkpoint_handler.save( self.directory, args=PyTreeSaveArgs(data, save_args=save_args_tree), diff --git a/checkpoint/orbax/checkpoint/_src/handlers/standard_checkpoint_handler_test.py b/checkpoint/orbax/checkpoint/_src/handlers/standard_checkpoint_handler_test.py index 08afba9b9..d25212b0f 100644 --- a/checkpoint/orbax/checkpoint/_src/handlers/standard_checkpoint_handler_test.py +++ b/checkpoint/orbax/checkpoint/_src/handlers/standard_checkpoint_handler_test.py @@ -329,9 +329,23 @@ def make_params(): ) test_utils.assert_tree_equal(self, params, restored) - def test_empty_error(self): + @parameterized.parameters( + (tuple([]),), + (dict(),), + (list(),), + ) + def test_empty_save(self, tree): + self.handler.save(self.directory, args=self.save_args_cls(tree)) + restored = self.handler.restore( + self.directory, args=self.restore_args_cls(tree) + ) + self.assertEqual(restored, tree) + + def test_none_save(self): + # Only `None` is unsaveable. Empty containers ({}, [], ()) are stored + # successfully as metadata-only entries at the root keypath. with self.assertRaises(ValueError): - self.handler.save(self.directory, args=self.save_args_cls({})) + self.handler.save(self.directory, args=self.save_args_cls(None)) def test_empty_dict_node(self): item = {'a': {}, 'b': 3} diff --git a/checkpoint/orbax/checkpoint/_src/metadata/tree.py b/checkpoint/orbax/checkpoint/_src/metadata/tree.py index 211f6ec38..4e399b6ec 100644 --- a/checkpoint/orbax/checkpoint/_src/metadata/tree.py +++ b/checkpoint/orbax/checkpoint/_src/metadata/tree.py @@ -919,13 +919,7 @@ def _is_bare_leaf(self, tree: PyTree) -> bool: A bare leaf is a checkpointable that is a single value rather than a container (e.g. ``save_pytree(dir, jnp.arange(5))``). Its lone leaf has the - empty keypath ``()``. This is distinct from an *empty* registered pytree - (0 leaves, e.g. an empty flax module), which is not a leaf and remains - unsupported here. - - None is explicitly excluded: JAX treats it as an empty pytree yet - treedef_is_leaf(tree_structure(None)) is True; None is not a valid - checkpointable and must not be held here. + empty keypath ``()``. Args: tree: A PyTree object to inspect. @@ -938,9 +932,17 @@ def _is_bare_leaf(self, tree: PyTree) -> bool: ) def _validate_tree_type(self, tree: PyTree): + """Validates that the tree type is supported.""" # Note: NamedTuple is a subclass of tuple. - if not isinstance(tree, (dict, list, tuple)) and not self._is_bare_leaf( - tree + # None is allowed as it represents an empty custom object when + # support_rich_types=False. An empty registered pytree (0 leaves, e.g. + # MyFlax) is allowed as it represents an empty custom object when + # support_rich_types=True. + if ( + tree is not None + and not isinstance(tree, (dict, list, tuple)) + and jax.tree.leaves(tree) + and not self._is_bare_leaf(tree) ): raise ValueError(f'Unsupported tree type: {type(tree)}') diff --git a/checkpoint/orbax/checkpoint/_src/metadata/tree_test.py b/checkpoint/orbax/checkpoint/_src/metadata/tree_test.py index 6734ebb17..b6ec4e423 100644 --- a/checkpoint/orbax/checkpoint/_src/metadata/tree_test.py +++ b/checkpoint/orbax/checkpoint/_src/metadata/tree_test.py @@ -336,14 +336,16 @@ def test_properties(self, tree): self._check_tree_property(tree, metadata) @parameterized.parameters( - # An empty registered pytree (0 leaves, not a container) is unsupported - # because it is neither a container nor a single leaf. + # An empty registered pytree (0 leaves, not a container) is supported + # as it represents an empty custom object when support_rich_types=True. (test_tree_utils.MyFlax(),), + # None is supported as it represents an empty custom object when + # support_rich_types=False. (None,), ) - def test_invalid_tree_type(self, tree): - with self.assertRaises(ValueError): - _TreeMetadataImpl(tree=tree) + def test_valid_empty_tree_type(self, tree): + metadata = _TreeMetadataImpl(tree=tree) + self.assertEqual(metadata.tree, tree) @parameterized.parameters( (1,), diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/handlers/pytree_handler_test.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/handlers/pytree_handler_test.py index 80d535e9e..74bc7c9d2 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/handlers/pytree_handler_test.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/handlers/pytree_handler_test.py @@ -1033,6 +1033,7 @@ def make_state_with_nones(): [], [1, [], 2], {'a': [], 'b': 3}, + (), ), array_metadata_store=(None, ARRAY_METADATA_STORE), ) @@ -1045,8 +1046,11 @@ def test_empty_data( with handler_with_options( use_ocdbt=use_ocdbt, array_metadata_store=array_metadata_store ) as checkpoint_handler: - if not data: - with self.assertRaisesRegex(ValueError, 'Found empty item'): + # Only `None` is unsaveable. Empty containers ({}, [], ()) are stored as + # a metadata-only entry at the root keypath and restore to the same + # empty container. + if data is None: + with self.assertRaisesRegex(ValueError, 'None is not saveable'): checkpoint_handler.save( self.directory, data, @@ -1064,6 +1068,31 @@ def test_empty_data( array_metadata_store=array_metadata_store, ) + @parameterized.product( + use_ocdbt=(True, False), + array_metadata_store=(None, ARRAY_METADATA_STORE), + ) + def test_empty_custom_object( + self, + use_ocdbt: bool, + array_metadata_store: array_metadata_store_lib.Store | None, + ): + """Tests saving and restoring an empty custom object like optax.EmptyState.""" + data = optax.EmptyState() + with handler_with_options( + use_ocdbt=use_ocdbt, array_metadata_store=array_metadata_store + ) as checkpoint_handler: + checkpoint_handler.save(self.directory, data) + # When loaded without a target item structure and + # support_rich_types=False, a custom empty object restores as None because + # it was assigned typestr='None'. + restored = checkpoint_handler.load(self.directory) + self.assertIsNone(restored) + + # When loaded with a target item structure, it restores perfectly. + restored_with_item = checkpoint_handler.load(self.directory, data) + self.assertEqual(restored_with_item, data) + @parameterized.product( use_ocdbt=(True, False), array_metadata_store=(None, ARRAY_METADATA_STORE), @@ -1517,8 +1546,14 @@ class PyTreeDict(dict): lambda keys, values: PyTreeDict(dict(zip(keys, values))), ) - with self.assertRaisesRegex(ValueError, 'Found empty item'): - self.handler.save(self.directory, PyTreeDict()) + # A top-level empty custom node saves as a metadata-only entry and, like + # the nested case below, restores as a plain dict (the custom container + # type is not preserved for empty nodes). + top_level_dir = self.directory / 'top_level_empty' + top_level_dir.mkdir(parents=True, exist_ok=True) + self.handler.save(top_level_dir, PyTreeDict()) + restored = self.handler.load(top_level_dir) # pylint: disable=g-unsafe-pickle-load + self.assertDictEqual({}, restored) self.handler.save(self.directory, {'a': PyTreeDict()}) restored = self.handler.load(self.directory) diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/partial/saving_test.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/partial/saving_test.py index 106f0bc5a..fc98fa194 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/partial/saving_test.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/partial/saving_test.py @@ -301,12 +301,16 @@ def get_dir_size(path: epath.Path) -> int: self.assertIn('arr2', restored) # pyrefly: ignore[bad-argument-type] def test_empty_initial_save(self): - """Tests that save() raises an error if the initial save is empty.""" + """Tests that save() is functional if the initial save is empty.""" final_path = self.directory / 'empty_initial_save' # First save - empty. - with self.assertRaisesRegex(ValueError, 'Found empty item.'): - saving.save(final_path, {}) # pyrefly: ignore[bad-argument-type] + saving.save(final_path, {}) # pyrefly: ignore[bad-argument-type] + saving.finalize(final_path) + self.assertTrue(final_path.exists()) + + restored_pytree = loading.load(final_path) + self.assertDictEqual({}, restored_pytree) @parameterized.named_parameters( ('none_then_meta', None, {'meta1': 'val1'}, {'meta1': 'val1'}), diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/testing/save_load_test.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/testing/save_load_test.py index 711e6e7bc..c670ff2fe 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/testing/save_load_test.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/testing/save_load_test.py @@ -234,11 +234,27 @@ async def mock_finalize(self_handler, directory): (tuple([]),), (dict(),), (list(),), + ) + def test_empty_native_tree(self, tree): + ocp.save(self.directory, tree) + with self.subTest('with_item'): + loaded = ocp.load(self.directory, tree) + self.assertEqual(tree, loaded) + with self.subTest('without_item'): + loaded = ocp.load(self.directory) + self.assertEqual(tree, loaded) + + @parameterized.parameters( (optax.EmptyState(),), ) - def test_empty_tree(self, tree): - with self.assertRaisesRegex(ValueError, 'Found empty item'): - ocp.save(self.directory, tree) + def test_empty_custom_node(self, custom_node): + ocp.save(self.directory, custom_node) + with self.subTest('with_item'): + loaded = ocp.load(self.directory, custom_node) + self.assertEqual(custom_node, loaded) + with self.subTest('without_item'): + loaded = ocp.load(self.directory) + self.assertIsNone(loaded) def test_none_tree(self): with self.assertRaisesRegex(