diff --git a/src/maxtext/common/grain_utility.py b/src/maxtext/common/grain_utility.py index 320048c3e5..fbd7e8ccc7 100644 --- a/src/maxtext/common/grain_utility.py +++ b/src/maxtext/common/grain_utility.py @@ -26,18 +26,301 @@ import jax from maxtext.input_pipeline import multihost_dataloading import numpy as np -import orbax.checkpoint as ocp -from orbax.checkpoint import v1 as ocp_v1 - +import orbax.checkpoint as ocp_v0 +from orbax.checkpoint import v1 as ocp PyGrainCheckpointHandler = python.PyGrainCheckpointHandler -Composite = ocp.args.Composite +Composite = ocp_v0.args.Composite RemoteIteratorWrapper = multihost_dataloading.RemoteIteratorWrapper ElasticIterator = grain_experimental.ElasticIterator +_PROCESS_0_FILENAME = "process_0.json" + + +def _process_filename(process_index: int, process_count: int) -> str: + return f"process_{process_index}-of-{process_count}.json" + + +def _state_to_text(iterator) -> str: + """Serializes an iterator's state (grain DatasetIterator -> json, else bytes).""" + if isinstance(iterator, grain.DatasetIterator): + return json.dumps(iterator.get_state(), indent=4) + return iterator.get_state().decode() + + +def _set_state_from_text(iterator, text: str) -> None: + """Inverse of :func:`_state_to_text`.""" + if isinstance(iterator, grain.DatasetIterator): + iterator.set_state(json.loads(text)) + else: + iterator.set_state(text.encode()) + + +async def no_op(): + return None + + +class GrainCheckpointable_v1(ocp.StatefulCheckpointable): + """Orbax v1 `StatefulCheckpointable` for MaxText grain data iterators. + + This is the v1 port of the `GrainCheckpointHandler`: a single object that + dispatches on the held item, kept here (rather than in grain) only for the + cases grain itself does not cover. The case-dispatch lives in this class — + callers just wrap the item — matching the original handler. + + `save` snapshots iterator state synchronously and returns a coroutine that + does the file IO in the background (the Orbax two-phase contract); `load` + returns a coroutine that reads and re-applies the state. + + Cases: + * **standard** grain iterator -> delegate to grain's own + `StatefulCheckpointable` (`process_{index}-of-{count}.json`). This + functionality is folded into grain. + * **ElasticIterator** -> one shared `process_0.json` (state is a single + global scalar; the fixed name survives a change in `jax.process_count()`). + Lift target: grain, which owns ``ElasticIterator``. + * **list** of `(iterator, process_index, process_count)` -> explicit + per-file write/read for a host-count change (`expansion_factor_real_data`). + The index/count arithmetic is computed by the caller. + * **RemoteIteratorWrapper** (Pathways colocated-python) -> the wrapper + persists/restores its own state keyed by `step`; identified structurally + (it exposes `save_state`/`restore_state`) so this module stays + independent of the input pipeline. + """ + + def __init__(self, item, *, restore_process_index=None, restore_process_count=None, step=None): + """Initializes a GrainCheckpointable_v1. + + Args: + item: a grain iterator, a grain `ElasticIterator`, a + `RemoteIteratorWrapper`, or a list of `(iterator, process_index, + process_count)` (scaled save). + restore_process_index: restore-only; an int (scale-up) or list of ints + (scale-down, paired with the items in `item`). `None` uses the + current process index. + restore_process_count: restore-only; the stored host count for the file name. + `None` uses the current process count (the standard, grain-native + path). + step: required only for the `RemoteIteratorWrapper` case. + """ + self._item = item + self._restore_process_index = restore_process_index + self._restore_process_count = restore_process_count + self._step = step + + async def save(self, directory): + """Snapshots the wrapped iterator's state and returns a coroutine that writes it to ``directory``.""" + item = self._item + + # RemoteIteratorWrapper handles checkpointing via colocated python + if isinstance(item, RemoteIteratorWrapper): + item.save_state(self._step) + return no_op() + + # ElasticIterator state is a single global scalar shared by all shards, + # so we write one fixed `process_0.json` from process 0 only. This file + # layout survives changes in `jax.process_count()`. + if isinstance(item, ElasticIterator): + state = item.get_state() # snapshot in the blocking phase + + async def _write_elastic(): + path = await directory.await_creation() + if jax.process_index() == 0: # one shared file written by process 0 + await asyncio.to_thread( + (path / _PROCESS_0_FILENAME).write_text, + json.dumps(state, indent=4), + ) + + return _write_elastic() + + if isinstance(item, list): + snapshots = [(_state_to_text(it), idx, count) for it, idx, count in item] + + async def _write_list(): + path = await directory.await_creation() + for text, idx, count in snapshots: + await asyncio.to_thread((path / _process_filename(idx, count)).write_text, text) + + return _write_list() + + if hasattr(item, "save"): + # Standard: delegate to grain's own StatefulCheckpointable. + return await item.save(directory) + + # Custom fallback for iterators without .save() + state_text = _state_to_text(item) + + async def _write_single_fallback(): + path = await directory.await_creation() + filename = path / _process_filename(jax.process_index(), jax.process_count()) + await asyncio.to_thread(filename.write_text, state_text) + + return _write_single_fallback() + + async def load(self, directory): + """Restores the wrapped iterator's state from ``directory`` (or returns a coroutine that does).""" + item = self._item + + # In Pathways + colocated_python environment, RemoteIteratorWrapper handles checkpointing + if isinstance(item, RemoteIteratorWrapper): + item.restore_state(self._step) + return no_op() + + # McJax and Pathways through controller cases + # ElasticIterator: every process reads the same shared `process_0.json`. + if isinstance(item, ElasticIterator): + + async def _read_elastic(): + filename = directory / _PROCESS_0_FILENAME + if not await asyncio.to_thread(filename.exists): + raise ValueError(f"File {filename} does not exist.") + item.set_state(json.loads(await asyncio.to_thread(filename.read_text))) + + return _read_elastic() + + if isinstance(item, list): + # Scale-down: each held iterator reads its own stored shard file. + specs = list(zip(item, self._restore_process_index)) # pyrefly: ignore[bad-argument-type] + + async def _read_list(): + for iterator, idx in specs: + filename = directory / _process_filename(idx, self._restore_process_count) # pyrefly: ignore[bad-argument-type] + if not await asyncio.to_thread(filename.exists): + raise ValueError(f"File {filename} does not exist.") + _set_state_from_text(iterator, await asyncio.to_thread(filename.read_text)) + + return _read_list() + + if self._restore_process_count is not None: + # Single iterator with an explicit stored host count (scale-up). + index = self._restore_process_index if self._restore_process_index is not None else jax.process_index() + + async def _read_single(): + filename = directory / _process_filename(index, self._restore_process_count) + if not await asyncio.to_thread(filename.exists): + raise ValueError(f"File {filename} does not exist.") + _set_state_from_text(item, await asyncio.to_thread(filename.read_text)) + + return _read_single() + + if hasattr(item, "load"): + # Standard: delegate to grain's own StatefulCheckpointable. + return await item.load(directory) + + # Custom fallback for iterators without .load() + async def _read_single_fallback(): + filename = directory / _process_filename(jax.process_index(), jax.process_count()) + if not await asyncio.to_thread(filename.exists): + raise ValueError(f"File {filename} does not exist.") + _set_state_from_text(item, await asyncio.to_thread(filename.read_text)) + + return _read_single_fallback() + + +def for_save(step: int, data_iterator: Any, expansion_factor_real_data: int) -> GrainCheckpointable_v1: + """Builds the v1 ``GrainCheckpointable_v1`` for saving the grain iterator.""" + if isinstance(data_iterator, RemoteIteratorWrapper): + return GrainCheckpointable_v1(data_iterator, step=step) + + if ( + not isinstance(data_iterator, list) + and hasattr(data_iterator, "local_iterator") + and isinstance(data_iterator.local_iterator, ElasticIterator) + ): + return GrainCheckpointable_v1(data_iterator.local_iterator) + + iterators = data_iterator if isinstance(data_iterator, list) else [data_iterator] + process_count_total = jax.process_count() * len(iterators) + if expansion_factor_real_data > 1: + process_count_total = process_count_total // expansion_factor_real_data + + if len(iterators) == 1 and process_count_total == jax.process_count(): + return GrainCheckpointable_v1( + iterators[0].local_iterator if hasattr(iterators[0], "local_iterator") else iterators[0] + ) + + specs = [ + ( + di.local_iterator if hasattr(di, "local_iterator") else di, + jax.process_index() + i * jax.process_count(), + process_count_total, + ) + for i, di in enumerate(iterators) + ] + return GrainCheckpointable_v1(specs) + + +def for_restore( + checkpoint_manager: Any, step: int, data_iterator: Any, expansion_factor_real_data: int +) -> GrainCheckpointable_v1: + """Builds the v1 ``GrainCheckpointable_v1`` for restoring the grain iterator.""" + if isinstance(data_iterator, RemoteIteratorWrapper): + return GrainCheckpointable_v1(data_iterator, step=step) + + if ( + not isinstance(data_iterator, list) + and hasattr(data_iterator, "local_iterator") + and isinstance(data_iterator.local_iterator, ElasticIterator) + ): + return GrainCheckpointable_v1(data_iterator.local_iterator) + + directory = checkpoint_manager.directory / str(step) / "iter" + process_count_jax = jax.process_count() + process_count_stored = len(list(directory.glob("process_*-of-*.json"))) + + if process_count_stored > process_count_jax: + assert isinstance(data_iterator, list), ( + f"{process_count_stored} processes found in Grain checkpoint directory {directory}, but only " + f"{process_count_jax} jax processes in this run, please set expansion_factor_real_data accordingly." + ) + scaling_factor = len(data_iterator) + expected = process_count_stored / process_count_jax + assert scaling_factor == expected, ( + f"Found {process_count_stored} processes in checkpoint and {process_count_jax} JAX processes, " + f"implying a scaling factor of {expected}, but the data_iterator list has {scaling_factor} items." + ) + local_iterators = [x.local_iterator if hasattr(x, "local_iterator") else x for x in data_iterator] + restore_process_index = [jax.process_index() + i * process_count_jax for i in range(scaling_factor)] + return GrainCheckpointable_v1( + local_iterators, restore_process_index=restore_process_index, restore_process_count=process_count_stored + ) + + if process_count_stored == process_count_jax: + assert not isinstance(data_iterator, list), ( + f"{process_count_stored} processes found in Grain checkpoint directory {directory}, matching the number of " + "jax processes, please do not set expansion_factor_real_data." + ) + return GrainCheckpointable_v1( + data_iterator.local_iterator if hasattr(data_iterator, "local_iterator") else data_iterator + ) + + if expansion_factor_real_data > 1 and process_count_stored == process_count_jax // expansion_factor_real_data: + assert not isinstance( + data_iterator, list + ), "when expansion_factor_real_data > 1, the data iterator should not be a list." + return GrainCheckpointable_v1( + data_iterator.local_iterator if hasattr(data_iterator, "local_iterator") else data_iterator, + restore_process_index=jax.process_index(), + restore_process_count=process_count_stored, + ) + + raise ValueError( + f"Error restoring Grain checkpoint in {directory}: " + f"The number of stored checkpoint files ({process_count_stored}) " + f"is incompatible with the number of JAX processes ({process_count_jax}). " + "If you are resuming training with a different number of chips, see instructions in " + "https://github.com/AI-Hypercomputer/maxtext/blob/main/docs/guides/data_input_pipeline/" + "data_input_grain.md#using-grain" + ) + + +# ------------------------------------------------------------------------------ +# TODO(b/532274266): Remove everything below this line once distillation_utils +# supports the new GrainCheckpointHandler. +# ------------------------------------------------------------------------------ -class GrainCheckpointHandler(PyGrainCheckpointHandler, ocp.CheckpointHandler): +class GrainCheckpointHandler(PyGrainCheckpointHandler, ocp_v0.CheckpointHandler): """A CheckpointHandler that allows specifying process_index and process_count.""" def save( @@ -131,21 +414,21 @@ def restore_single_process(item, process_index, process_count): return restore_single_process(item, process_index, process_count) -@ocp.args.register_with_handler(GrainCheckpointHandler, for_save=True) +@ocp_v0.args.register_with_handler(GrainCheckpointHandler, for_save=True) @dataclasses.dataclass -class GrainCheckpointSave(ocp.args.CheckpointArgs): +class GrainCheckpointSave(ocp_v0.args.CheckpointArgs): item: Any -@ocp.args.register_with_handler(GrainCheckpointHandler, for_restore=True) +@ocp_v0.args.register_with_handler(GrainCheckpointHandler, for_restore=True) @dataclasses.dataclass -class GrainCheckpointRestore(ocp.args.CheckpointArgs): +class GrainCheckpointRestore(ocp_v0.args.CheckpointArgs): item: Any process_index: Optional[int | list[int]] = None process_count: Optional[int] = None -class GrainCheckpointable(ocp_v1.StatefulCheckpointable): +class GrainCheckpointable(ocp.StatefulCheckpointable): """Adapts `GrainCheckpointHandler` to Orbax v1's `StatefulCheckpointable`.""" def __init__( diff --git a/tests/unit/grain_utility_test.py b/tests/unit/grain_utility_test.py new file mode 100644 index 0000000000..1dd42a4180 --- /dev/null +++ b/tests/unit/grain_utility_test.py @@ -0,0 +1,343 @@ +# 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 for the consolidated grain v1 ``GrainCheckpointable``.""" + +import asyncio +import json +import pathlib +import tempfile +from typing import Any +import unittest +from unittest import mock + +from absl.testing import absltest +from absl.testing import parameterized +from etils import epath +import grain +from grain import experimental as grain_experimental +import grain.sharding +from maxtext.common import grain_utility +from orbax.checkpoint import v1 as ocp + + +ElasticIterator = grain_experimental.ElasticIterator +GrainCheckpointable = grain_utility.GrainCheckpointable_v1 + + +def _std_iter() -> Any: + return iter(grain.MapDataset.range(10).to_iter_dataset()) + + +def _elastic_iter(): + return ElasticIterator( + grain.MapDataset.range(40), + global_batch_size=2, + shard_options=grain.sharding.ShardOptions(shard_index=0, shard_count=1), + ) + + +class _FakeRemote(grain_utility.RemoteIteratorWrapper): + """Stands in for RemoteIteratorWrapper (colocated-python save/restore by step).""" + + def __init__(self): + # pylint: disable=super-init-not-called + self.saved_step = None + self.restored_step = None + + def save_state(self, step): + self.saved_step = step + + def restore_state(self, step): + self.restored_step = step + + +async def _drive(stateful_coro): + background = await stateful_coro + await background + + +class TestGrainCheckpointable(unittest.TestCase): + """Tests for the GrainCheckpointable class.""" + + def test_standard_delegates_to_grain_native(self): + with tempfile.TemporaryDirectory() as d: + path = pathlib.Path(d) / "ckpt" + it = _std_iter() + for _ in range(4): + next(it) + expected = it.get_state() + ocp.save_checkpointables(str(path), {"iter": GrainCheckpointable(it)}) + # grain-native writes its own per-process file name. + self.assertTrue((path / "iter" / "process_0-of-1.json").exists()) + restored = _std_iter() + ocp.load_checkpointables(str(path), {"iter": GrainCheckpointable(restored)}) + self.assertEqual(restored.get_state(), expected) + + def test_elastic_writes_single_shared_file(self): + with tempfile.TemporaryDirectory() as d: + path = pathlib.Path(d) / "ckpt" + it = _elastic_iter() + for _ in range(3): + next(it) + expected = it.get_state() + ocp.save_checkpointables(str(path), {"iter": GrainCheckpointable(it)}) + self.assertTrue((path / "iter" / "process_0.json").exists()) # reshard-safe + restored = _elastic_iter() + ocp.load_checkpointables(str(path), {"iter": GrainCheckpointable(restored)}) + self.assertEqual(restored.get_state(), expected) + + def test_scaled_list_explicit_index_count(self): + with tempfile.TemporaryDirectory() as d: + path = pathlib.Path(d) / "ckpt" + a, b = _std_iter(), _std_iter() + for _ in range(2): + next(a) + for _ in range(6): + next(b) + state_a, state_b = a.get_state(), b.get_state() + ocp.save_checkpointables(str(path), {"iter": GrainCheckpointable([(a, 0, 2), (b, 1, 2)])}) + self.assertTrue((path / "iter" / "process_0-of-2.json").exists()) + self.assertTrue((path / "iter" / "process_1-of-2.json").exists()) + ra, rb = _std_iter(), _std_iter() + ocp.load_checkpointables( + str(path), + {"iter": GrainCheckpointable([ra, rb], restore_process_index=[0, 1], restore_process_count=2)}, + ) + self.assertEqual(ra.get_state(), state_a) + self.assertEqual(rb.get_state(), state_b) + + def test_scale_up_single_with_explicit_count(self): + with tempfile.TemporaryDirectory() as d: + path = pathlib.Path(d) / "ckpt" + it = _std_iter() + for _ in range(7): + next(it) + expected = it.get_state() + ocp.save_checkpointables(str(path), {"iter": GrainCheckpointable([(it, 0, 1)])}) + restored = _std_iter() + ocp.load_checkpointables( + str(path), {"iter": GrainCheckpointable(restored, restore_process_index=0, restore_process_count=1)} + ) + self.assertEqual(restored.get_state(), expected) + + def test_remote_forwards_step(self): + wrapper = _FakeRemote() + asyncio.run(_drive(GrainCheckpointable(wrapper, step=7).save(None))) + asyncio.run(_drive(GrainCheckpointable(wrapper, step=9).load(None))) + self.assertEqual(wrapper.saved_step, 7) + self.assertEqual(wrapper.restored_step, 9) + + +# pylint: disable=protected-access +class TestGrainUtilityFactories(unittest.TestCase): + """Tests for the high-level factories for_save and for_restore.""" + + def test_for_save_remote(self): + wrapper = _FakeRemote() + c = grain_utility.for_save(7, wrapper, 1) + self.assertEqual(c._step, 7) + self.assertIs(c._item, wrapper) + + def test_for_save_elastic(self): + it = _elastic_iter() + # Mock data_iterator to have .local_iterator + mock_iter = mock.MagicMock() + mock_iter.local_iterator = it + c = grain_utility.for_save(0, mock_iter, 1) + self.assertIs(c._item, it) + + def test_for_save_standard(self): + it = _std_iter() + c = grain_utility.for_save(0, it, 1) + self.assertIs(c._item, it) + + def test_for_save_scaled_list(self): + a, b = _std_iter(), _std_iter() + c = grain_utility.for_save(0, [a, b], 1) + self.assertIsInstance(c._item, list) + self.assertEqual(len(c._item), 2) + # specs: (item, index, total) + self.assertIs(c._item[0][0], a) + self.assertEqual(c._item[0][1], 0) + self.assertEqual(c._item[0][2], 2) + + def test_for_restore_remote(self): + wrapper = _FakeRemote() + c = grain_utility.for_restore(None, 7, wrapper, 1) + self.assertEqual(c._step, 7) + self.assertIs(c._item, wrapper) + + def test_for_restore_elastic(self): + it = _elastic_iter() + mock_iter = mock.MagicMock() + mock_iter.local_iterator = it + c = grain_utility.for_restore(None, 0, mock_iter, 1) + self.assertIs(c._item, it) + + +# pylint: enable=protected-access + + +# ------------------------------------------------------------------------------ +# TODO(b/532274266): Remove everything below this line once distillation_utils +# supports the new GrainCheckpointHandler. +# ------------------------------------------------------------------------------ + + +class GrainCheckpointableEquivalenceTest(parameterized.TestCase): + """Tests to ensure GrainCheckpointable is equivalent to GrainCheckpointHandler.""" + + def setUp(self): + super().setUp() + self.tmp_dir = epath.Path(self.create_tempdir().full_path) + + def test_save_restore_equivalence_single_item(self): + class FakeIterator: + """A fake iterator for testing.""" + + def __init__(self, state=0): + self.state = state + + def get_state(self): + return json.dumps({"state": self.state}).encode() + + def set_state(self, state): + self.state = json.loads(state.decode())["state"] + + def __next__(self): + self.state += 1 + return self.state + + iterator_v0 = FakeIterator(10) + iterator_v1 = FakeIterator(10) + + step = 100 + v0_path = self.tmp_dir / str(step) / "iter_v0" + v1_path = self.tmp_dir / str(step) / "iter_v1" + + # v0 Save + handler = grain_utility.GrainCheckpointHandler() + v0_path.mkdir(parents=True, exist_ok=True) + handler.save(v0_path, item=iterator_v0) + + # v1 Save + wrapper = GrainCheckpointable(iterator_v1) + + class MockDirectory: + + async def await_creation(self): + v1_path.mkdir(parents=True, exist_ok=True) + return v1_path + + commit_func = asyncio.run(wrapper.save(MockDirectory())) + if commit_func: + asyncio.run(commit_func) + + # Verify files are identical + v0_file = v0_path / "process_0-of-1.json" + v1_file = v1_path / "process_0-of-1.json" + + self.assertTrue(v0_file.exists()) + self.assertTrue(v1_file.exists()) + self.assertEqual(v0_file.read_text(), v1_file.read_text()) + + # v0 Restore + restored_iterator_v0 = FakeIterator(0) + args_v0 = grain_utility.GrainCheckpointRestore(item=restored_iterator_v0) + handler.restore(v0_path, args=args_v0) + self.assertEqual(restored_iterator_v0.state, 10) + + # v1 Restore + restored_iterator_v1 = FakeIterator(0) + wrapper_restore = GrainCheckpointable(restored_iterator_v1) + + load_func = asyncio.run(wrapper_restore.load(v1_path)) + asyncio.run(load_func) + self.assertEqual(restored_iterator_v1.state, 10) + + def test_save_restore_equivalence_list_item(self): + class FakeIterator: + """A fake iterator for testing.""" + + def __init__(self, state=0): + self.state = state + + def get_state(self): + return json.dumps({"state": self.state}).encode() + + def set_state(self, state): + self.state = json.loads(state.decode())["state"] + + iterator_a = FakeIterator(10) + iterator_b = FakeIterator(20) + + item_v0 = [(iterator_a, 0, 2), (iterator_b, 1, 2)] + item_v1 = [(iterator_a, 0, 2), (iterator_b, 1, 2)] + + step = 100 + v0_path = self.tmp_dir / str(step) / "iter_v0" + v1_path = self.tmp_dir / str(step) / "iter_v1" + + # v0 Save + handler = grain_utility.GrainCheckpointHandler() + v0_path.mkdir(parents=True, exist_ok=True) + handler.save(v0_path, item=item_v0) + + # v1 Save + wrapper = GrainCheckpointable(item_v1) + + class MockDirectory: + + async def await_creation(self): + v1_path.mkdir(parents=True, exist_ok=True) + return v1_path + + commit_func = asyncio.run(wrapper.save(MockDirectory())) + if commit_func: + asyncio.run(commit_func) + + # Verify files are identical + v0_file_0 = v0_path / "process_0-of-2.json" + v1_file_0 = v1_path / "process_0-of-2.json" + v0_file_1 = v0_path / "process_1-of-2.json" + v1_file_1 = v1_path / "process_1-of-2.json" + + self.assertTrue(v0_file_0.exists()) + self.assertTrue(v1_file_0.exists()) + self.assertEqual(v0_file_0.read_text(), v1_file_0.read_text()) + + self.assertTrue(v0_file_1.exists()) + self.assertTrue(v1_file_1.exists()) + self.assertEqual(v0_file_1.read_text(), v1_file_1.read_text()) + + # v0 Restore + iterators_restore_v0 = [FakeIterator(0), FakeIterator(0)] + args_v0 = grain_utility.GrainCheckpointRestore(item=iterators_restore_v0, process_index=[0, 1], process_count=2) + handler.restore(v0_path, args=args_v0) + + self.assertEqual(iterators_restore_v0[0].state, 10) + self.assertEqual(iterators_restore_v0[1].state, 20) + + # v1 Restore + iterators_restore_v1 = [FakeIterator(0), FakeIterator(0)] + wrapper_restore = GrainCheckpointable(iterators_restore_v1, restore_process_index=[0, 1], restore_process_count=2) + load_func = asyncio.run(wrapper_restore.load(v1_path)) + asyncio.run(load_func) + self.assertEqual(iterators_restore_v1[0].state, 10) + self.assertEqual(iterators_restore_v1[1].state, 20) + + +if __name__ == "__main__": + absltest.main()