| 494 | """ |
| 495 | |
| 496 | def fn(value: Union[Tensor, VDict]) -> NestedTensor: |
| 497 | if not isinstance(value, VDict): |
| 498 | return value |
| 499 | |
| 500 | leaves = jax.tree_util.tree_leaves(value) |
| 501 | if not leaves: |
| 502 | # An empty VDict. |
| 503 | return value |
| 504 | |
| 505 | non_tensor_leaves = [leaf for leaf in leaves if not isinstance(leaf, Tensor)] |
| 506 | if non_tensor_leaves: |
| 507 | raise ValueError( |
| 508 | f"Expected a tree of Tensors, got {type(non_tensor_leaves[0])} in {tree}" |
| 509 | ) |
| 510 | |
| 511 | scalar_tensors = [leaf for leaf in leaves if not leaf.shape] |
| 512 | if scalar_tensors: |
| 513 | raise ValueError( |
| 514 | f"Expected a tree of vectorized Tensors, got scalar {scalar_tensors} in {tree}" |
| 515 | ) |
| 516 | |
| 517 | vdict_size = leaves[0].shape[0] |
| 518 | different_vdict_size_tensors = [leaf for leaf in leaves if leaf.shape[0] != vdict_size] |
| 519 | if different_vdict_size_tensors: |
| 520 | raise ValueError( |
| 521 | "Expected a tree of vectorized Tensors of same dim 0, " |
| 522 | f"got {different_vdict_size_tensors[0].shape[0]} vs. {vdict_size} in {tree}" |
| 523 | ) |
| 524 | |
| 525 | expanded: list[VDict] = [] |
| 526 | for ind in range(vdict_size): |
| 527 | value_i: VDict = jax.tree.map(lambda x, i=ind: x[i], value) |
| 528 | expanded_i = {k: expand_vdicts(v) for k, v in value_i.items()} |
| 529 | expanded.append(expanded_i) |
| 530 | return expanded |
| 531 | |
| 532 | return jax.tree.map(fn, tree, is_leaf=lambda x: isinstance(x, VDict)) |
| 533 | |