已开启
【RFC】声明式并行编程支持动态图张量并行 #9
yangzhenzhang创建于 1月26日
1月26日 修改了issue 的描述
此处折叠了26条事件消息 查看更多
4月16日 关联了pull request:Update the inference logic and unit tests of some distributed operators
4月16日 关联了pull request:Update the inference logic and unit tests of some distributed operators
背景
张量并行是大模型训练中最基础的技术之一,是将模型中各算子的输入/输出张量在多卡上进行切分的技术。
当前业界大模型训练大量使用Megatron的方案,将模型的拆分、通信以分布式单元的形式构造出来。Megatron方案的好处是显而易见的,所有分布式的行为都展示给了用户,用户可自主控制的范围很大,比如用户可以在网络各个位置进行修改及调试;
但Megatron方案的也有严重的弊端,即并行逻辑与模型深度耦合。当引入新模态或调整模型结构时,常常需要深入底层、重写大量的分布式代码。同时,算法研究与工程团队需紧密绑定,而不能快速、独立地验证不同想法。
如今,业界各知名的分布式训练框架也在构建与算法模型解耦的张量并行能力。比如,Pytorch围绕Dtensor对张量进行切分间模,同时字节也基于Torch的Dtensor推出了vesclae框架,Jax/OneFlow也有类似的方案。
Hyper Parallel提供了DTensor,基于DTensor来构建张量并行能力。
方案
基本原理
算子级并行即将算子的输入张量切分到多张卡上,使得每张卡计算一个分片,其汇总的结果是整个算子的计算结果。
神经网络结构可以简单划分为三个部分:正向、反向、优化器。1)正向:MindSpore对正向部分的每个算子及其输入张量进行切分建模,使得该算子的计算逻辑在切分前后保持数学等价,正向部分的每个算子都有对应的分布式算子实现。2)反向:反向算子没有对应的分布式算子实现,正向切分完成后,反向使用自动微分能力;3)优化器:优化器算子也没有对应的分布式算子实现,在优化器定义之前,传入的是切片权重,使得优化器内部的权重及优化器状态都是分片,直接使用这些分片进行计算。
在算子级并行中,有两个关键的建模,分布式张量分布(Tensor Layout)和张量重排布(Tensor Redistribution)。张量分布表示张量切分后,张量切片在集群分布情况。张量重排布表示两种张量分布之间的转换。模型切分流程如下图所示,首先对每个算子的输入张量按策略进行切分,生成算子输入的张量分布,然后根据算子的数学定义,推导出输出张量分布;然后再检查前一个算子输出张量分布和下一个算子的输入张量分布,如果两种不同,则会插入张量重排布。
算子1 --> 输出张量排布 --> 张量重排布 --> 输入张量排布 --> 算子2
另外,反向起始的sens需要有两方面的处理:1)需要根据loss的切分形态对sens也做同样的切分,否则反向算子的输入shape将不匹配;2)如果loss重复计算N份,为了保证最终梯度计算结果与单卡等价,sens的值需要除以N;
算子捕获入口
张量并行需要对每个算子进行建模,获取算子输入的切分信息,从而进行处理。算子的输入张量需为DTensor,我们在DTensor中构建机制来捕获对应的算子:
使用mindspore提供的__fallback__机制,它能截获正向网络的算子,算子粒度相当于mindspore的Primitive。
样例如下:
使用torch提供的__torch_function__机制,它能截获正向网络的算子,算子粒度相当于torch的function接口。
捕获了算子之后,便可以对算子进行策略传播、算子替换等并行处理;
策略传播
策略传播即根据算子输入切分策略推导其输出切分策略,其原理为:
1,捕获到算子;
2,提取算子每个输入张量的切分策略及算子行为强相关信息(比如matmul的transpore_a/transpose_b, reshape的目的shape等);
3,根据这些信息,推导出算子输出的切分策略;
注:当前还不支持框架内部自动策略搜索;
算子切分类型
算子切分大致可划分为三种类型:
可直接使用原算子及其所有输入
输入张量的切分,并没有破坏原来算子的计算规则,没有其他非张量输入或其他非张量输入也不破坏原来算子的计算规则。
可直接使用原算子及所有输入进行计算,比如relu算子;
需要更改算子输入
非张量输入需随着张量输入的切分而更改,比如reshape算子;
例:单卡逻辑out = reshape(input, (8, 8))
假设input的行切分8份,那么需要将目的shape修改为(1, 8)
算子的输入张量切分后,原算子的计算规则被破坏,导致无法直接使用原算子进行计算,需要将其替换成其他算子的组合;
例:单卡逻辑out = linear(input, weight, bias)
假设input的列以及weight的行被切了8份,那么使用原linear算子对分片进行计算会导致数学不等价。
此时,一种方式是将linear换成 matmul + allreduce + biasadd的算子组合。
相关接口
注:接口正在评审刷新中
1,shard(model, sharding_plan):
作用:网络cell或函数配置并行切分策略
输入:
model:cell或函数
sharding_plan:字典,用于描述并行切分策略,内部包含两个部分:
1)“parameter”:用于描述cell内部的所有权重如何切分;
2)“forward”:用于描述cell及子cell的输入/输出如何切分;
3)如果model是函数,则不需要“parameter”部分;
为cell配置sharding_plan的示例如下:
sharding_plan = {
"parameter" : {
"weight" : layout0, # 顶层cell的权重切分策略配置
"bias" : layout1,
"sub_net.dense.weight": layout2, # 子cell的权重切分策略配置
"sub_net.dense.bias": layout3,
},
"forward" : {
"input": {"x": layout4}, # 顶层输入的切分策略配置,示例配置为dict类型,还支持list类型:如“input”:[layout4]
"output":{"out": layout5}}, # 顶层输出的切分策略配置,示例配置为dict类型,还支持list类型
"sub_net.input": {"x": layout6}, # 子cell输入的切分策略配置,示例配置为dict类型,还支持list类型
"sub_net.ouput": {"out", layout7} # 子cell输出的切分策略配置,示例配置为dict类型,还支持list类型
}
}
输出:如果model是cell,则返回cell;如果model是函数,则返回封装后的函数;
2,parallelize_value_and_grad(fn, weights, sens=None):
作用:用于分布式训练过程的求导,它在内部会对sens进行并行处理,如做切分及根据正向的输出排布自动调整反向sens值;
输入:
fn:cell或函数
weights:训练网络的权重
sens:(Union[list(float), tuple(float)], optional), 反向求导的输入,默认是None;
1)如果fn只有一个输出,那么sens必须为None,它在函数内部自动处理;
2)如果fn有多个输出:若sens为None,则内部让第一个sens为1.0,其余为0;若sens不为None,它的长度必须和fn的输出长度相同;
返回:
封装后的函数(调用封装后的函数,将返回fn的输出及权重的梯度列表)
使用约束
测试设计
1)构造一个单卡用例,用例中的网络使用已支持张量并行的算子;
2)使用同一网络,构造一个并行用例;
3)分别跑一个训练step,观察正向loss/反向grad是否能完全对齐;