Flattens `tree` and returns a list of (path, value) pairs.
(
tree: Nested[Tensor], separator: str = "/", is_leaf: Optional[Callable[[Any], bool]] = None
)
| 419 | |
| 420 | |
| 421 | def flatten_items( |
| 422 | tree: Nested[Tensor], separator: str = "/", is_leaf: Optional[Callable[[Any], bool]] = None |
| 423 | ) -> Sequence[tuple[str, Tensor]]: |
| 424 | """Flattens `tree` and returns a list of (path, value) pairs.""" |
| 425 | flat_paths_and_values, _ = jax.tree_util.tree_flatten_with_path(tree, is_leaf=is_leaf) |
| 426 | return list( |
| 427 | (separator.join(_key_entry_to_str(k) for k in path), value) |
| 428 | for path, value in flat_paths_and_values |
| 429 | ) |
| 430 | |
| 431 | |
| 432 | @jax.tree_util.register_pytree_with_keys_class |