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)
| 1052 | |
| 1053 | |
| 1054 | def 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, |