# Source code for tensorplay.distributed.device_mesh Source: https://www.tensorplay.cn/docs/_modules/tensorplay/distributed/device_mesh.html ``` # # plain sizes/strides bookkeeping; the public API (init_device_mesh, # get_group, get_local_rank, get_coordinate, submesh __getitem__, context # manager, from_group) is preserved. import math import threading from typing import Any import tensorplay as tp import tensorplay.distributed as dist __all__ = ["DeviceMesh", "init_device_mesh"] class _MeshNDimensionality(int): def __new__(cls, value: int) -> "_MeshNDimensionality": return int.__new__(cls, value) def __call__(self) -> int: return int(self) class _MeshEnv(threading.local): def __init__(self) -> None: # root_mesh_to_mesh self.root_to_flat_mesh: dict[Any, dict[str, Any]] = {} self.mesh_stack: list[Any] = [] @staticmethod def get() -> "_MeshEnv": if not hasattr(_MeshEnv, "_local"): _MeshEnv._local = _MeshEnv() return _MeshEnv._local class _MeshResources: def __init__(self) -> None: # map from root_mesh to list of all meshes with root_mesh as parent self.root_to_2d_mesh: dict = {} def create_sub_mesh( self, root_mesh, sub_mesh, mesh_dim_names ) -> None: self.root_to_2d_mesh.setdefault(root_mesh, {})[mesh_dim_names] = sub_mesh _mesh_resources = _MeshResources() def _get_device_handle(device_type: str = "cuda"): return getattr(tp, device_type, None) def _flatten(sizes): out = [] strides = [1] * len(sizes) acc = 1 for i in reversed(range(len(sizes))): strides[i] = acc acc *= sizes[i] for idx in range(math.prod(sizes)): coords = [] rem = idx for s in sizes: rem, c = divmod(rem, s) coords.append(c) coords.reverse() flat = sum(c * st for c, st in zip(coords, strides)) out.append(flat) return out def _coord_at(sizes, strides, flat_idx): if flat_idx < 0 or flat_idx >= math.prod(sizes): raise IndexError("flat mesh index is out of range") rem = int(flat_idx) coords = [] for stride, size in zip(strides, sizes): coord, rem = divmod(rem, stride) if coord >= size: raise IndexError("flat mesh index is out of range") coords.append(coord) return tuple(coords) [docs] class DeviceMesh: """ The mesh is an n-d array whose values are global ranks. Process groups are created per mesh dimension so collectives can run on each dimension independently. Example:: >>> from tensorplay.distributed.device_mesh import init_device_mesh >>> mesh = init_device_mesh("cuda", mesh_shape=(2, 4), ... mesh_dim_names=("dp", "tp")) """ def __init__( self, device_type: str, mesh=None, *, mesh_dim_names=None, _dim_group_names=None, _rank_map=None, _sizes=None, _strides=None, _root_mesh=None, _backend_override=None, _axis_root_dims=None, ): if mesh is not None and (_rank_map is not None or _sizes is not None): raise TypeError( "Cannot provide internal fields when passing an explicit mesh" ) if mesh is not None: if isinstance(mesh, tp.Tensor): mesh = mesh.cpu().tolist() if isinstance(mesh, int): mesh = [mesh] def flatten(value): if not isinstance(value, (list, tuple)): return (), [int(value)] if not value: return (0,), [] child_shapes = [] flat_values = [] for child in value: child_shape, child_values = flatten(child) child_shapes.append(child_shape) flat_values.extend(child_values) if any(shape != child_shapes[0] for shape in child_shapes[1:]): raise ValueError("all mesh dimensions must be rectangular") return (len(value),) + child_shapes[0], flat_values sizes, rank_map = flatten(mesh) if not sizes: sizes = (1,) else: if _rank_map is None or _sizes is None: raise TypeError("The mesh argument is required") rank_map = list(_rank_map) sizes = list(_sizes) if not sizes or any( isinstance(size, bool) or not isinstance(size, int) or size <= 0 for size in sizes ): raise ValueError("mesh dimensions must be positive integers") sizes = tuple(sizes) total = math.prod(sizes) if len(rank_map) != total: raise AssertionError( f"rank map length {len(rank_map)} != product of sizes {total}" ) if any(not isinstance(rank, int) or rank < 0 for rank in rank_map): raise ValueError("mesh ranks must be non-negative integers") if len(set(rank_map)) != len(rank_map): raise ValueError("mesh ranks must be unique") # row-major strides strides = [1] * len(sizes) acc = 1 for i in reversed(range(len(sizes))): strides[i] = acc acc *= sizes[i] if mesh_dim_names is not None: if len(set(mesh_dim_names)) != len(mesh_dim_names): raise ValueError("Each mesh_dim_name must be unique.") if len(mesh_dim_names) != len(sizes): raise ValueError( "mesh_shape and mesh_dim_names should have same length!" ) self._mesh_dim_names = tuple(mesh_dim_names) else: self._mesh_dim_names = None self.device_type = device_type self._sizes = tuple(sizes) self._strides = tuple(strides) self._rank_map = list(rank_map) def nest(values, shape): if len(shape) == 1: return list(values) width = math.prod(shape[1:]) return [nest(values[i * width:(i + 1) * width], shape[1:]) for i in range(shape[0])] self.mesh = nest(self._rank_map, self._sizes) self._root_mesh = _root_mesh self._thread_id: int | None = None self._flatten_mapping: dict[str, "DeviceMesh"] = {} self._dim_groups: dict[int, Any] = {} self._backend_override = _backend_override if _axis_root_dims is not None: if len(_axis_root_dims) != len(self._sizes): raise ValueError("axis metadata must match mesh dimensions") self._axis_root_dims = tuple( tuple(int(dim) for dim in dims) for dims in _axis_root_dims ) else: root = self._get_root_mesh() root_names = getattr(root, "_mesh_dim_names", None) axis_root_dims = [] for index, name in enumerate(self._mesh_dim_names or ()): if root_names is not None and name in root_names: axis_root_dims.append((root_names.index(name),)) else: axis_root_dims.append((index,)) if len(axis_root_dims) != len(self._sizes): axis_root_dims = [(index,) for index in range(len(self._sizes))] self._axis_root_dims = tuple(axis_root_dims) if _dim_group_names is not None: self._dim_group_names = list(_dim_group_names) else: self._dim_group_names = self._init_process_groups() # ------------------------------------------------------------------ # process-group setup # ------------------------------------------------------------------ def _my_rank(self) -> int: try: return dist.get_rank() except Exception: return -1 def _coords_of(self, pos: int) -> tuple[int, ...]: coords = [] rem = pos for d, st in zip(self._sizes, self._strides): coords.append(rem // st) rem %= st return tuple(coords) def _pos_of_coords(self, coords) -> int: if len(coords) != len(self._sizes): raise ValueError("coordinate rank does not match mesh dimensions") if any(coord < 0 or coord >= size for coord, size in zip(coords, self._sizes)): raise IndexError("mesh coordinate is out of range") return sum(c * s for c, s in zip(coords, self._strides)) def _ranks_along_dim(self, dim: int, coords: tuple[int, ...]) -> list[int]: """Return ranks on the mesh line containing ``coords``.""" from itertools import product ranges = [ range(size) if axis == dim else (coords[axis],) for axis, size in enumerate(self._sizes) ] return [ self._rank_map[self._pos_of_coords(combo)] for combo in product(*ranges) ] def _init_process_groups(self) -> list[str]: names = [] for dim in range(len(self._sizes)): name = (self._mesh_dim_names[dim] if self._mesh_dim_names else f"dim_{dim}") names.append(name) # Defer actual subgroup creation until get_group is called; the # per-dim groups are created lazily by the SPMD ranks together. self._lazy_groups = True return names def _get_or_create_group_for_dim(self, mesh_dim) : """Create (once) the subgroup along `mesh_dim` containing this rank.""" if isinstance(mesh_dim, str): if self._mesh_dim_names is None: raise KeyError(mesh_dim) try: dim = self._mesh_dim_names.index(mesh_dim) except ValueError as exc: raise KeyError(mesh_dim) from exc else: if isinstance(mesh_dim, bool): raise TypeError("mesh_dim must be an integer or string") dim = int(mesh_dim) if dim < 0 or dim >= self.ndim(): raise IndexError("mesh dimension is out of range") if dim in self._dim_groups: return self._dim_groups[dim] my_global = dist.get_rank() try: pos = self._rank_map.index(my_global) except ValueError as e: raise RuntimeError( f"Rank {my_global} is not part of this DeviceMesh" ) from e coords = self._coords_of(pos) group = None from itertools import product other_ranges = [ range(size) if axis != dim else (0,) for axis, size in enumerate(self._sizes) ] for line_coords in product(*other_ranges): ranks = self._ranks_along_dim(dim, tuple(line_coords)) kwargs = {"ranks": ranks} if self._backend_override is not None: kwargs["backend"] = self._backend_override candidate = dist.new_group(**kwargs) if my_global in ranks: group = candidate if group is None: raise RuntimeError(f"Rank {my_global} is not part of mesh dimension {dim}") self._dim_groups[dim] = group return group # ------------------------------------------------------------------ # public API # ------------------------------------------------------------------ def get_group(self, mesh_dim=None): if mesh_dim is None: if int(self.ndim) == 1: mesh_dim = 0 else: raise RuntimeError( "Must specify mesh_dim for multi-dimensional mesh." ) return self._get_or_create_group_for_dim(mesh_dim) def size(self, mesh_dim: int | None = None) -> int: if mesh_dim is None: return len(self._rank_map) if isinstance(mesh_dim, str): if self._mesh_dim_names is None: raise KeyError(mesh_dim) mesh_dim = self._mesh_dim_names.index(mesh_dim) mesh_dim = int(mesh_dim) if mesh_dim < 0 or mesh_dim >= self.ndim(): raise IndexError("mesh dimension is out of range") return self._sizes[mesh_dim] @property def ndim(self) -> int: return _MeshNDimensionality(len(self._sizes)) @property def shape(self) -> tuple[int, ...]: return self._sizes def numel(self) -> int: return math.prod(self._sizes) def get_rank(self) -> int: return dist.get_rank() def get_all_groups(self) -> list[Any]: return [self.get_group(index) for index in range(int(self.ndim))] @property def mesh_dim_names(self): return self._mesh_dim_names @property def ndimension(self) -> int: return len(self._sizes) def get_local_rank(self, mesh_dim=None) -> int: try: my_global = dist.get_rank() except RuntimeError: my_global = 0 try: pos = self._rank_map.index(my_global) except ValueError as exc: raise RuntimeError( f"Rank {my_global} is not part of this DeviceMesh" ) from exc coords = self._coords_of(pos) if mesh_dim is None: if int(self.ndim) != 1: raise RuntimeError("Must specify mesh_dim.") mesh_dim = 0 if isinstance(mesh_dim, str): if self._mesh_dim_names is None: raise KeyError(mesh_dim) dim = self._mesh_dim_names.index(mesh_dim) else: dim = int(mesh_dim) if dim < 0 or dim >= self.ndim(): raise IndexError("mesh dimension is out of range") return coords[dim] [docs] def get_coordinate(self) -> tuple[int, ...] | None: """Returns this rank's coordinate in the mesh, or None if absent.""" try: my_global = dist.get_rank() except RuntimeError: my_global = 0 try: pos = self._rank_map.index(my_global) except ValueError: return None return self._coords_of(pos) def __getitem__(self, mesh_dim_names) -> "DeviceMesh": if isinstance(mesh_dim_names, str): mesh_dim_names = (mesh_dim_names,) if self._mesh_dim_names is None: raise RuntimeError( "No `mesh_dim_names` found; cannot slice the mesh." ) if not mesh_dim_names: raise ValueError("at least one mesh dimension must be selected") if len(set(mesh_dim_names)) != len(mesh_dim_names): raise ValueError("mesh dimensions must be unique") if tuple(mesh_dim_names) == self._mesh_dim_names: return self try: dims = tuple(self._mesh_dim_names.index(n) for n in mesh_dim_names) except ValueError as exc: root = self._get_root_mesh() if len(mesh_dim_names) == 1 and mesh_dim_names[0] in root._flatten_mapping: return root._flatten_mapping[mesh_dim_names[0]] raise KeyError(mesh_dim_names) from exc sub_sizes = tuple(self._sizes[d] for d in dims) coordinate = self.get_coordinate() if coordinate is None: coordinate = tuple(0 for _ in self._sizes) # A slice over a subset of dimensions is the line through this rank; # selecting every dimension retains the complete mesh. selected_ranges = [range(self._sizes[d]) for d in dims] from itertools import product rank_map = [] for sel in product(*selected_ranges): dim_to_coord = dict(zip(dims, sel)) coords = [dim_to_coord.get(d, coordinate[d]) for d in range(len(self._sizes))] rank_map.append(self._rank_map[self._pos_of_coords(coords)]) sub = DeviceMesh( device_type=self.device_type, _rank_map=rank_map, _sizes=sub_sizes, mesh_dim_names=mesh_dim_names, _root_mesh=self._get_root_mesh(), _backend_override=self._backend_override, _axis_root_dims=tuple( self._get_axis_root_dims()[dim] for dim in dims ), ) _mesh_resources.create_sub_mesh( self._get_root_mesh(), sub, mesh_dim_names ) return sub def _flatten(self, mesh_dim_name: str | None = None, backend_override=None) -> "DeviceMesh": if not self._mesh_dim_names: raise RuntimeError("Cannot flatten a mesh without dimension names") if mesh_dim_name is None: mesh_dim_name = "_".join(self._mesh_dim_names) if not isinstance(mesh_dim_name, str) or not mesh_dim_name: raise ValueError("flattened mesh dimension name must be non-empty") if self.ndim == 1 and mesh_dim_name == self._mesh_dim_names[0]: return self root = self._get_root_mesh() if root._mesh_dim_names and mesh_dim_name in root._mesh_dim_names: raise ValueError( f"{mesh_dim_name} already exists in the root mesh dimensions" ) existing = root._flatten_mapping.get(mesh_dim_name) if existing is not None: if existing._rank_map != self._rank_map or existing._sizes != (len(self._rank_map),): raise ValueError( f"flattened mesh dimension {mesh_dim_name!r} already has a different layout" ) return existing flattened = DeviceMesh( self.device_type, list(self._rank_map), mesh_dim_names=(mesh_dim_name,), _root_mesh=root, _backend_override=( backend_override if backend_override is not None else self._backend_override ), _axis_root_dims=( tuple( dim for axis in self._get_axis_root_dims() for dim in axis ), ), ) root._flatten_mapping[mesh_dim_name] = flattened return flattened def _get_axis_root_dims(self) -> tuple[tuple[int, ...], ...]: value = getattr(self, "_axis_root_dims", None) if value is not None: return value root = self._get_root_mesh() root_names = getattr(root, "_mesh_dim_names", None) result = [] for index, name in enumerate(self._mesh_dim_names or ()): if root_names is not None and name in root_names: result.append((root_names.index(name),)) else: result.append((index,)) if len(result) != len(self._sizes): result = [(index,) for index in range(len(self._sizes))] self._axis_root_dims = tuple(result) return self._axis_root_dims @staticmethod def _concatenate(device_mesh_list: list["DeviceMesh"]) -> "DeviceMesh": if not device_mesh_list: raise ValueError("at least one DeviceMesh is required") first = device_mesh_list[0] root = first._get_root_mesh() if any(not isinstance(mesh, DeviceMesh) for mesh in device_mesh_list): raise TypeError("all entries must be DeviceMesh instances") if any(mesh._get_root_mesh() is not root for mesh in device_mesh_list): raise RuntimeError( "Cannot concatenate DeviceMeshes derived from different device meshes" ) if any(mesh.device_type != first.device_type for mesh in device_mesh_list): raise RuntimeError("Cannot concatenate DeviceMeshes with different device types") names: list[str] = [] axes: list[tuple[int, ...]] = [] for mesh in device_mesh_list: mesh_names = mesh.mesh_dim_names if mesh_names is None or len(mesh_names) != int(mesh.ndim): raise ValueError("all DeviceMeshes must have mesh dimension names") mesh_axes = mesh._get_axis_root_dims() if len(mesh_axes) != len(mesh_names): raise ValueError("mesh axis metadata is invalid") names.extend(str(name) for name in mesh_names) axes.extend(mesh_axes) root_sizes = root._sizes used_dims: set[int] = set() axis_sizes: list[int] = [] for axis in axes: if not axis or any(dim < 0 or dim >= len(root_sizes) for dim in axis): raise RuntimeError("Cannot concatenate invalid mesh axes") if used_dims.intersection(axis): raise RuntimeError( f"Cannot concatenate overlapping meshes: {device_mesh_list}" ) used_dims.update(axis) axis_sizes.append(math.prod(root_sizes[dim] for dim in axis)) coordinate = root.get_coordinate() if coordinate is None: coordinate = tuple(0 for _ in root_sizes) rank_map: list[int] = [] from itertools import product for axis_coordinates in product(*(range(size) for size in axis_sizes)): full_coordinate = list(coordinate) for axis, axis_coordinate in zip(axes, axis_coordinates): remaining = int(axis_coordinate) for index, root_dim in enumerate(axis): inner_size = math.prod( root_sizes[other_dim] for other_dim in axis[index + 1 :] ) value = remaining // inner_size remaining %= inner_size full_coordinate[root_dim] = value rank_map.append(root._rank_map[root._pos_of_coords(full_coordinate)]) result = DeviceMesh( first.device_type, _rank_map=rank_map, _sizes=tuple(axis_sizes), mesh_dim_names=tuple(names), _root_mesh=root, _backend_override=first._backend_override, _axis_root_dims=tuple(axes), ) output_dim = 0 for mesh in device_mesh_list: for dim in range(int(mesh.ndim)): if dim in mesh._dim_groups: result._dim_groups[output_dim] = mesh._dim_groups[dim] output_dim += 1 return result def _get_root_mesh(self) -> "DeviceMesh": return self._root_mesh if self._root_mesh else self def __enter__(self) -> "DeviceMesh": _MeshEnv.get().mesh_stack.append(self) return self def __exit__(self, exc_type, exc_value, exc_traceback) -> None: _MeshEnv.get().mesh_stack.pop() def __repr__(self) -> str: if self._mesh_dim_names: dims_repr = ", ".join( f"{k}={v}" for k, v in zip(self._mesh_dim_names, self._sizes) ) else: dims_repr = str(tuple(self._sizes)) return f"DeviceMesh({dims_repr}, '{self.device_type}')" def __eq__(self, other: object) -> bool: if self is other: return True if not isinstance(other, DeviceMesh): return False return ( self._rank_map == other._rank_map and self._sizes == other._sizes and self.device_type == other.device_type and self._mesh_dim_names == other._mesh_dim_names ) def __hash__(self): return hash(( tuple(self._rank_map), tuple(self._sizes), self.device_type, self._mesh_dim_names, )) def __getstate__(self) -> dict[str, Any]: state = dict(self.__dict__) state["_dim_groups"] = {} return state def __setstate__(self, state: dict[str, Any]) -> None: self.__dict__.update(state) self._dim_groups = {} [docs] @classmethod def from_group(cls, group, device_type=None, mesh=None, mesh_dim_names=None) -> "DeviceMesh": """Construct a DeviceMesh from one or more existing process groups.""" device_type = device_type or "cuda" groups = list(group) if isinstance(group, (list, tuple)) else [group] if not groups: raise ValueError("at least one process group is required") if mesh_dim_names is not None and len(mesh_dim_names) != len(groups): raise ValueError("mesh_dim_names must match the number of groups") group_ranks = [dist.get_process_group_ranks(item) for item in groups] if len(groups) == 1: ranks = group_ranks[0] if mesh is None: mesh = ranks else: if isinstance(mesh, tp.Tensor): mesh = mesh.cpu().tolist() flat = [] def collect(value): if isinstance(value, (list, tuple)): for item in value: collect(item) else: flat.append(int(value)) collect(mesh) if flat != ranks: raise ValueError("mesh must list process-group ranks in order") elif mesh is None: raise ValueError("mesh is required when multiple groups are provided") result = cls(device_type=device_type, mesh=mesh, mesh_dim_names=mesh_dim_names) result._dim_groups = {index: item for index, item in enumerate(groups)} return result [docs] def init_device_mesh( device_type: str, mesh_shape: tuple[int, ...], *, mesh_dim_names: tuple[str, ...] | None = None, backend_override=None, ) -> DeviceMesh: """ This creates a DeviceMesh with an n-dimensional array layout, where `n` is the length of `mesh_shape`. If `mesh_dim_names` is provided, each dimension is labeled as `mesh_dim_names[i]`. .. note:: Follows SPMD: ensure `mesh_shape` is identical across all ranks. Example:: >>> mesh_1d = init_device_mesh("cuda", mesh_shape=(8,)) >>> mesh_2d = init_device_mesh("cuda", mesh_shape=(2, 8), ... mesh_dim_names=("dp", "tp")) """ if mesh_dim_names is not None: if len(set(mesh_dim_names)) != len(mesh_dim_names): raise RuntimeError( "Each mesh_dim_name must be unique. " f"Found repeated mesh_dim_name in mesh_dim_names {mesh_dim_names}" ) if len(mesh_shape) != len(mesh_dim_names): raise RuntimeError( "mesh_shape and mesh_dim_names should have same length! " f"Found len(mesh_dim_names): {len(mesh_dim_names)} and " f"len(mesh_shape):{len(mesh_shape)}." ) if not isinstance(device_type, str) or not device_type or not device_type.isalpha(): raise RuntimeError( f"Device type with index is not supported but got {device_type}. ", "If you maintained a 'tp.device' object, it's recommended to " "pass in 'device.type'.", ) if not dist.is_initialized(): raise RuntimeError( "init_device_mesh requires tensorplay.distributed to be " "initialized first (call dist.init_process_group)." ) world_size = dist.get_world_size() if not mesh_shape or any( isinstance(size, bool) or not isinstance(size, int) or size <= 0 for size in mesh_shape ): raise ValueError("mesh_shape must contain positive integers") mesh_size = math.prod(mesh_shape) if mesh_size != world_size: raise RuntimeError( f"mesh_shape product ({mesh_size}) must equal world size ({world_size})" ) return DeviceMesh( device_type=device_type, mesh=_reshape_ranks(list(range(world_size)), mesh_shape), mesh_dim_names=mesh_dim_names, _backend_override=backend_override, ) def _reshape_ranks(ranks: list[int], shape: tuple[int, ...]): """Row-major reshape of a flat rank list into nested lists (n-d).""" if len(shape) == 1: return ranks if len(shape) == 2: r, c = shape return [ranks[i * c : (i + 1) * c] for i in range(r)] # generic n-d nesting outer = shape[0] chunk = len(ranks) // outer rest = shape[1:] return [_reshape_ranks(ranks[i * chunk : (i + 1) * chunk], rest) for i in range(outer)] ```