API symbolsdistributed
Copy
DeviceMesh
- class tensorplay.distributed.device_mesh.DeviceMesh(device_type: str, mesh=None, *, mesh_dim_names=None, _dim_group_names=None, _rank_map=None, _sizes=None, _strides=None, _root_mesh=None)[source]
DeviceMesh represents a mesh of devices (torch parity).
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"))
- classmethod from_group(group, device_type=None, mesh=None, mesh_dim_names=None) DeviceMesh[source]
Construct a 1-D DeviceMesh from an existing ProcessGroup.
- get_coordinate() tuple[int, ...] | None[source]
Returns this rank’s coordinate in the mesh, or None if absent.
- get_group(mesh_dim=None)[source]
Returns the process group along
mesh_dim(torch parity).
Help improve this page
Found an error, an unclear step, or a missing example?
Was this page helpful?
