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

Method tree_flatten_with_keys

axlearn/common/checkpointer_test.py:1568–1577  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

1566
1567 def __repr__(self):
1568 return SWITCHABLE_VDICT_IMPL.__repr__(self)
1569
1570 def tree_flatten_with_keys(self):
1571 try:
1572 return SWITCHABLE_VDICT_IMPL.tree_flatten_with_keys(self)
1573 except NotImplementedError:
1574 # tree_paths() no longer generates named keys for pytrees that don't register
1575 # their child keys, so we simulate the child keys from the old implementation
1576 # of tree_paths().
1577 values, keys = SWITCHABLE_VDICT_IMPL.tree_flatten(self)
1578 key_value = [(jax.tree_util.DictKey(k), v) for k, v in zip(keys, values)]
1579 return key_value, keys
1580

Callers

nothing calls this directly

Calls 1

tree_flattenMethod · 0.45

Tested by

no test coverage detected