已开启
【RFC】DTensor及张量切分表达 #10
yangzhenzhang创建于 1月26日
1月26日 修改了issue 的描述
1月26日 修改了issue 的描述
1月26日 修改了issue 的描述
1月27日 修改了issue 的描述
2月2日 关闭了 issue
2月2日 issue状态由 TODO 改变为 DONE
2月3日 issue状态由 DONE 改变为 WIP
2月3日 重新打开了 issue
2月14日 关闭了 issue
2月14日 issue状态由 WIP 改变为 DONE
2月26日 issue状态由 DONE 改变为 WIP
2月26日 重新打开了 issue
3月27日 关联了pull request:fix: remove pipeline stage check from repeat_num
Yyangzhenzhang
6月16日 关联了pull request:test(dtensor): add unit tests to improve coverage for core/dtensor modules
6月16日 关联了pull request:test(dtensor): add unit tests to improve coverage for core/dtensor modules
Yyangzhenzhang
6月22日 关联了pull request:fix: use contextvars for SkipDTensorDispatch to support safe nesting
6月22日 关联了pull request:fix: use contextvars for SkipDTensorDispatch to support safe nesting
背景
目前,扩展分布式训练主要有三种方式:数据并行(Data Parallel)、张量并行(Tensor Parallel)和流水线并行(Pipeline Parallel),每种方式都独立地工作在不同的维度上。在训练非常大的模型时,用户希望能够将这些技术结合起来使用(即三维并行)。理想的情况是,用户可以像在单个节点/设备上一样构建模型,而无需担心如何在集群中进行分布式训练,为了实现这一目标,我们需要将单个设备的模型转换为分布式版本,并使用我们的运行时进行训练和部署。因此分布式并行框架需要有一套机制来统一的表达张量的排布方式,用于描述张量如何分片、哪个分片在哪些卡上切分,在哪些卡上重复等。
目前业界主要有两种表达风格,一种是torch的风格,另一种是jax的风格。
torch风格可参考:https://docs.pytorch.org/docs/stable/distributed.tensor.html
torch主要用DeviceMesh和Placement来提供统一的张量分片布局。其中DeviceMesh 表示设备拓扑,Placement 描述 DeviceMesh 上逻辑张量的分片布局。
jax风格可参考:https://docs.jax.dev/en/latest/notebooks/Distributed_arrays_and_automatic_parallelization.html
二者的主要区别在于:PyTorch采用面向网格维度的视角,即对于设备网格中的每个维度,你指定应该应用哪种切分方式;而JAX则采用面向张量维度的视角,即对于张量上的每个维度,你指定它应该被哪个网格维度(可能多个!)来切分。
方案
我们当前在接口层面使用DeviceMesh 和 placement的表达方式,而内部将其转换成类Jax的风格。
DeviceMesh
DeviceMesh主要包含两个信息:设备矩阵及别名;
创建示例:
它表示由编号为0/1/2/3的四张卡构成2 * 2 的设备矩阵,其中设备矩阵的行的别名为“dp”, 设备矩阵的列的别名为 “tp”。
支持从DeviceMesh中取出子mesh。
示例:
它表示从device_mesh中取出“dp”维度的子mesh。对rank0或rank1来说,dp_mesh的rank_list是[0, 1];而对rank2或rank3来说,dp_mesh的rank_list是[2, 3]
get_groupget_local_rankflattenget_device_num_along_axisget_rank_list_along_axisget_global_shapePlacement
Placement定义了设备矩阵每个维度对张量的切分方式。hyper_parallel当前提供了三种分布策略:
Shard、Replicate和Partial。通过DeviceMesh和Placement表达DTensor
示例:
为8卡创建一个4 * 2的设备矩阵,设备矩阵的行/列别名分别为“dp”和“tp”;

对于一个8 * 2的tensor来说,那么不同的placement将会让其有不同的切分,如图所示:
注:placement支持填入设备矩阵的轴别名,为了能处理单个轴切多次的情况;
内部处理逻辑
由于hyper-parallel原先使用jax风格描述张量切分形态,为了尽可能小的改动,内部模块依然还是用类jax风格的layout来描述。
因此,需要对Shard/Replicate/Partial的描述进行转换。
示例1:
处理逻辑为:
1,创建初始device-mesh,mesh_shape:(2, 2), alias_name: (“dp”, “tp”), rank-list:(0, 1, 2, 3)
2,获取子mesh,mesh_shape: (2),alias_name: (“dp”),rank_list: (0, 2) ——注:rank0/2视角
3,在from_local()内部:
1)使用Layout.from_device_mesh()创建layout实例;
2)调用layout((Shard(0),)),记录placement的信息;
3)调用layout. placement_to_tensor_map(),传入local_x的维度,完成风格转换;
示例2:
处理逻辑为:
1,创建初始device-mesh,mesh_shape:(2, 2), alias_name: (“dp”, “tp”), rank-list:(0, 1, 2, 3)
2,获取子mesh,mesh_shape: (2),alias_name: (“dp”),rank_list: (0, 2) ——注:rank0/2视角
3,在from_local()内部:
1)使用Layout.from_device_mesh()创建layout实例;
2)调用layout(("None",))
注:在例1中,由于使用了Shard风格,内部需要调用layout. placement_to_tensor_map()进行转换;而在例2中,使用轴别名字符串,则不需要调用layout. placement_to_tensor_map()进行转换;
对外接口
注:接口还在持续评审刷新中
张量排布相关接口
init_device_mesh(mesh_shape, alias_name) -> DeviceMesh
1,参数说明:1)mesh_shape:设备矩阵的shape;2)alias_name:设备矩阵的轴别名;
2,作用:创建DeviceMesh,其中,设备卡号认为是0~设备矩阵总元素 - 1
class DeviceMesh(mesh, alias_name) -> DevcieMesh
1,参数说明:1)mesh:设备矩阵;2)alias_name:设备矩阵的轴别名;
2,作用:创建DeviceMesh
DeviceMesh的相关方法:
__getitem__get_groupget_local_rankflattenget_device_num_along_axisget_rank_list_along_axisget_global_shapeDTensor相关接口
DTensor.from_local(local_tensor, device_mesh, placements) -> DTensor
1,参数说明:1)local_tensor:本地张量切片;2)device_mesh:设备mesh;3)placements:切分方式;
2,作用:创建DTensor,其中local_tensor是本地张量切片,即切分后的样子;
distribute_tensor(full_tensor, device_mesh, placements) -> DTensor
1,参数说明:1)local_tensor:本地完整张量;2)device_mesh:设备mesh;3)placements:切分方式;
2,作用:创建DTensor,其中full_tensor是逻辑上的完整张量,即切分前的样子;
DTensor.redistribute(self, device_mesh, placements) -> DTensor
1,参数说明:1)self: 原始DTensor;2)device_mesh:设备mesh;3)placements:切分方式;
2,作用:将原始DTensor转换成期望的排布;
相关约束