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

Function pytree_children

axlearn/common/utils.py:1935–1957  ·  view source on GitHub ↗

Generate the (key, value) pairs for the immediate children of a pytree `node`. Reference: jax._src.tree_util.generate_key_paths() Example: ``` assert pytree_children(dict(a=[1,2])) == [(DictKey('a'), [1,2])] ```

(node: Any)

Source from the content-addressed store, hash-verified

1933 logging.warning("Falling back to ICI-only mesh on GPU, performance may be reduced.")
1934 return build_standard_mesh(mesh_shape, devices=devices)
1935
1936 # Canonicalize to HybridMeshShape. If DCN mesh is not specified, break the first non-singleton
1937 # device axis (the least communication intensive) over the number of slices/granules. If all
1938 # axes are singletons, this is effectively a no-op, since this implies a single-granule
1939 # environment.
1940 if isinstance(mesh_shape, MeshShape):
1941 mesh_shape = infer_mesh_shape(mesh_shape, num_devices=num_devices)
1942 for axis, dim in enumerate(mesh_shape):
1943 if dim % num_granules == 0:
1944 break
1945 elif dim != 1:
1946 raise ValueError(
1947 f"First non-singleton mesh axis {axis} with value {dim} must be divisible by "
1948 f"the number of slices/granules {num_granules}."
1949 )
1950 else:
1951 raise ValueError(
1952 f"At least one axis of {mesh_shape=} must be divisible by {num_granules=}."
1953 )
1954
1955 if num_granules > 1:
1956 logging.info("Building multi-slice/granule device mesh over axis %s.", axis)
1957 # Truncate intra-slice/granule mesh.
1958 mesh_shape = (*mesh_shape[:axis], dim // num_granules, *mesh_shape[axis + 1 :])
1959 logging.info("Inferred intra-slice/granule mesh shape: %s", mesh_shape)
1960 # Configure data center (inter-slice/granule) mesh.

Callers 2

test_pytree_childrenMethod · 0.90
_find_cyclesFunction · 0.85

Calls

no outgoing calls

Tested by 1

test_pytree_childrenMethod · 0.72