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)
| 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. |
no outgoing calls