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

Function validate_contains_paths

axlearn/common/utils.py:2056–2065  ·  view source on GitHub ↗

Raises ValueError if any of the given `paths` are not present in `x`.

(x: Nested[Tensor], paths: Sequence[str])

Source from the content-addressed store, hash-verified

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
2067def find_cycles(tree: Nested) -> dict[str, KeyPath]:
2068 """Find a cycle in pytree `tree` if one exists.

Callers 15

predictMethod · 0.90
forwardMethod · 0.90
beam_search_decodeMethod · 0.90
prefill_statesMethod · 0.90
extend_stepMethod · 0.90
beam_search_decodeMethod · 0.90
beam_search_decodeMethod · 0.90
sample_decodeMethod · 0.90
_forward_for_modeMethod · 0.90
forwardMethod · 0.90
prefill_statesMethod · 0.90
test_basicMethod · 0.90

Calls 1

get_recursivelyFunction · 0.85

Tested by 1

test_basicMethod · 0.72