(self, iterator: Iterable)
| 495 | self.assertEqual(num_batches, len(text_examples) // logical_feed_logical_batch_size) |
| 496 | |
| 497 | def _check_iterator_saveable(self, iterator: Iterable): |
| 498 | # Check that we can save the data iterator. |
| 499 | with tempfile.TemporaryDirectory() as td: |
| 500 | save_dir = os.path.join(td, "ckpt") |
| 501 | step = 100 |
| 502 | ckptr = ( |
| 503 | Checkpointer.default_config() |
| 504 | .set(name="ckptr", dir=save_dir) |
| 505 | .instantiate(parent=None) |
| 506 | ) |
| 507 | with Mesh(jax.devices(), "data"): |
| 508 | ckptr.save(step=step, state={"iterator": iterator}, evaler_summaries=None) |
| 509 | ckptr.wait_until_finished() |
| 510 | self.assertTrue(os.path.exists(os.path.join(save_dir, f"step_{step:08d}", "index"))) |
| 511 | |
| 512 | @parameterized.product( |
| 513 | num_physical_feeds=(2, 4), |
no test coverage detected