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

Function async_save_tf_savables

axlearn/common/checkpointer.py:203–229  ·  view source on GitHub ↗

Asynchronously saves TF savables from `value_map` into `dir`. When this call returns, `value_map` can be safely mutated, but saving to `dir` will not complete unless the returned future is set.

(
    value_map: Nested[Any], *, executor: futures.ThreadPoolExecutor, dir: str
)

Source from the content-addressed store, hash-verified

201
202# pylint: disable=redefined-builtin
203def async_save_tf_savables(
204 value_map: Nested[Any], *, executor: futures.ThreadPoolExecutor, dir: str
205) -> futures.Future:
206 """Asynchronously saves TF savables from `value_map` into `dir`.
207
208 When this call returns, `value_map` can be safely mutated, but saving to `dir` will not
209 complete unless the returned future is set.
210 """
211 # pylint: disable-next=consider-using-with
212 f = tempfile.TemporaryDirectory()
213 for path, value in utils.flatten_items(value_map):
214 if isinstance(value, tf.data.Iterator):
215 tf_checkpoint = tf.train.Checkpoint(value)
216 tf_checkpoint.write(os.path.join(f.name, f"tf_{jax.process_index()}", path))
217 elif isinstance(value, ElasticDatasetIterator):
218 primary_tf_checkpoint = tf.train.Checkpoint(value.primary_iterator)
219 primary_tf_checkpoint.write(os.path.join(f.name, f"tf_{jax.process_index()}", path))
220
221 if value.elastic_iterator is not None:
222 if value.is_primary_for_checkpoint:
223 for pid in value.elastic_process_ids:
224 elastic_tf_checkpoint = tf.train.Checkpoint(value.elastic_iterator)
225 elastic_tf_checkpoint.write(os.path.join(f.name, f"tf_{pid}", path))
226 else:
227 raise NotImplementedError(f"Unknown dataset iterator: {value}")
228
229 return executor.submit(_upload_dir, f, dst_dir=dir)
230
231
232# pylint: disable-next=redefined-builtin

Callers 3

save_to_dirMethod · 0.85

Calls 3

joinMethod · 0.80
writeMethod · 0.45
submitMethod · 0.45

Tested by 2