Raises ValueError if any of the given `paths` are not present in `x`.
(x: Nested[Tensor], paths: Sequence[str])
| 2054 | return [] |
| 2055 | |
| 2056 | if isinstance(node, tuple) and hasattr(node, "_fields") and flat[1] == type(node): # noqa: E721 |
| 2057 | # Handle namedtuple as a special case, based on heuristic. |
| 2058 | return [(jax.tree_util.GetAttrKey(s), getattr(node, s)) for s in node._fields] |
| 2059 | |
| 2060 | key_children, _ = jax.tree_util.default_registry.flatten_one_level_with_keys(node) |
| 2061 | if key_children: |
| 2062 | return key_children |
| 2063 | |
| 2064 | return [(jax.tree_util.FlattenedIndexKey(i), c) for i, c in enumerate(flat[0])] |
| 2065 | |
| 2066 | |
| 2067 | def find_cycles(tree: Nested) -> dict[str, KeyPath]: |
| 2068 | """Find a cycle in pytree `tree` if one exists. |