MCPcopy Create free account
hub / github.com/AI-Hypercomputer/maxtext / create_device_mesh

Function create_device_mesh

src/MaxText/maxtext_utils.py:1054–1124  ·  view source on GitHub ↗

Creates a device mesh with each slice in its own data parallel group. If there is only one slice, uses two replicas

(config, devices=None)

Source from the content-addressed store, hash-verified

1052
1053
1054def create_device_mesh(config, devices=None):
1055 """Creates a device mesh with each slice in its own data parallel group. If there is only one slice, uses two replicas"""
1056 if devices is None:
1057 devices = jax.devices()
1058 if config.subslice_shape and config.enable_single_controller and config.num_slices == 1:
1059 max_logging.log(f"Trying to create a subslice with shape: {config.subslice_shape}")
1060 subslice_shape = tuple(int(x) for x in config.subslice_shape.split(","))
1061 device_coords = [device.coords for device in devices]
1062 device_coords_np = np.array(device_coords)
1063
1064 # Find the minimum coordinates to start the subslice
1065 min_coords = device_coords_np.min(axis=0)
1066
1067 subslice_devices = []
1068 for device in devices:
1069 coords = device.coords
1070 if all(min_coords[i] <= coords[i] < min_coords[i] + subslice_shape[i] for i in range(len(subslice_shape))):
1071 subslice_devices.append(device)
1072 devices = subslice_devices
1073
1074 num_devices = len(devices)
1075 num_slices = 1 if config.inference_benchmark_test else config.num_slices
1076 num_devices_per_slice = num_devices // num_slices
1077
1078 multi_slice_env = num_slices > 1
1079
1080 # Find possible unspecified parallelisms
1081 ici_parallelism = max_utils.fill_unspecified_mesh_axes(config.ici_parallelism.copy(), num_devices_per_slice, "ICI")
1082
1083 allow_split_physical_axes = config.allow_split_physical_axes if config.allow_split_physical_axes else False
1084
1085 if multi_slice_env:
1086 dcn_parallelism = max_utils.fill_unspecified_mesh_axes(config.dcn_parallelism.copy(), num_slices, "DCN")
1087 if max_utils.is_valid_custom_mesh(ici_parallelism, config.custom_mesh):
1088 mesh = max_utils.create_custom_device_mesh(ici_parallelism, dcn_parallelism, devices, config.custom_mesh)
1089 else:
1090 mesh = mesh_utils.create_hybrid_device_mesh(
1091 ici_parallelism,
1092 dcn_parallelism,
1093 devices,
1094 allow_split_physical_axes=allow_split_physical_axes,
1095 )
1096 else:
1097 if allow_split_physical_axes:
1098 if max_utils.is_valid_custom_mesh(ici_parallelism, config.custom_mesh):
1099 mesh = mesh_utils.create_device_mesh(
1100 [16, 16],
1101 devices,
1102 contiguous_submeshes=False,
1103 allow_split_physical_axes=False,
1104 )
1105 mesh = max_utils.reshape_mesh_to_rings(mesh, config.custom_mesh)
1106 mesh = np.reshape(mesh, ici_parallelism)
1107 else:
1108 mesh = mesh_utils.create_device_mesh(
1109 ici_parallelism,
1110 devices,
1111 contiguous_submeshes=False,

Callers 2

setUpMethod · 0.90
setUpMethod · 0.90

Calls 1

copyMethod · 0.80

Tested by 2

setUpMethod · 0.72
setUpMethod · 0.72