(self, checkpointer_cls: Type[BaseCheckpointer])
| 1132 | self.assertEqual(checkpointer_cls.latest_checkpoint_path(td), final_ckpt_path) |
| 1133 | self.assertEqual(checkpointer_cls.latest_checkpoint_step(td), 10) |
| 1134 | |
| 1135 | @parameterized.parameters([Checkpointer, OrbaxCheckpointer]) |
| 1136 | def test_read_state_spec(self, checkpointer_cls: Type[BaseCheckpointer]): |
| 1137 | mesh_shape = (1, 1) |
| 1138 | if not test_utils.is_supported_mesh_shape(mesh_shape): |
| 1139 | return |
| 1140 | with _mesh(mesh_shape): |
| 1141 | cfg = _checkpointer_config(checkpointer_cls) |
| 1142 | cfg.save_policy.min_step = 0 |
| 1143 | ckpt: BaseCheckpointer = cfg.instantiate(parent=None) |
| 1144 | state0 = dict( |
| 1145 | **{ |
| 1146 | f"v_{str(dtype.dtype)}": jnp.zeros([], dtype=dtype) |
| 1147 | for dtype in (jnp.uint32, jnp.int32, jnp.int64) |
| 1148 | }, |
| 1149 | **{ |
| 1150 | f"v_{str(dtype.dtype)}": jnp.zeros([4], dtype=dtype) |
| 1151 | for dtype in (jnp.float16, jnp.float32, jnp.float64) |
| 1152 | }, |
| 1153 | **{ |
| 1154 | f"v_{str(dtype.dtype)}": jnp.zeros([4, 2], dtype=dtype) |
| 1155 | for dtype in (jnp.bfloat16, jnp.bool_) |
| 1156 | }, |
| 1157 | ) |
| 1158 | ckpt.save(step=0, state=state0) |
| 1159 | ckpt.wait_until_finished() |
| 1160 | # Tests `read_state_spec`. |
| 1161 | state_spec = read_state_spec(checkpointer_cls.latest_checkpoint_path(cfg.dir)) |
| 1162 | self.assertNestedEqual( |
| 1163 | state_spec, |
| 1164 | jax.tree.map(lambda t: utils.TensorSpec(shape=t.shape, dtype=t.dtype), state0), |
| 1165 | ) |
| 1166 | step, state1 = ckpt.restore(state=state_spec) |
| 1167 | self.assertNestedEqual(0, step) |
| 1168 | self.assertNestedEqual(state0, state1) |
| 1169 |
nothing calls this directly
no test coverage detected