MCPcopy Create free account
hub / github.com/apple/axlearn / _check_iterator_saveable

Method _check_iterator_saveable

axlearn/common/input_tf_data_test.py:497–510  ·  view source on GitHub ↗
(self, iterator: Iterable)

Source from the content-addressed store, hash-verified

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),

Calls 6

joinMethod · 0.80
instantiateMethod · 0.45
setMethod · 0.45
default_configMethod · 0.45
saveMethod · 0.45
wait_until_finishedMethod · 0.45

Tested by

no test coverage detected