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
)
| 201 | |
| 202 | # pylint: disable=redefined-builtin |
| 203 | def 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 |