已开启
【RFC】声明式并行编程支持动态图张量并行 #9
yangzhenzhang创建于  1月26日
yangzhenzhang成员
1月26日 创建

背景

张量并行是大模型训练中最基础的技术之一,是将模型中各算子的输入/输出张量在多卡上进行切分的技术。

当前业界大模型训练大量使用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平台
    使用mindspore提供的__fallback__机制,它能截获正向网络的算子,算子粒度相当于mindspore的Primitive。
    样例如下:
def __fallback__(self, func, args={}, kwargs=None):
        if kwargs is None:
            kwargs = {}
        from hyper_parallel.core.shard._op_dispatch import _OP_DISPATCHER
        with NoFallbackGuard():
            out = _OP_DISPATCHER.dispatch(func, args, kwargs)
        return out
  • torch平台
    使用torch提供的__torch_function__机制,它能截获正向网络的算子,算子粒度相当于torch的function接口。
    @classmethod
    def __torch_function__(
        cls,
        func: torch._C._FunctionBase,
        types: Tuple[type, ...],
        args: Tuple[Any, ...] = (),
        kwargs: Optional[Dict[str, Any]] = None
    ) -> Any:
        kwargs = kwargs or {}
        from hyper_parallel.core.tensor_parallel._op_dispatch import _OP_DISPATCHER
        out = _OP_DISPATCHER.dispatch(func, args, kwargs)
        return out

捕获了算子之后,便可以对算子进行策略传播、算子替换等并行处理;

策略传播

策略传播即根据算子输入切分策略推导其输出切分策略,其原理为:
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的输出及权重的梯度列表)

使用约束

  • 网络正向使用到的算子需支持分布式(除非使用custom_shard()封装的部分);
  • 对于训练,需结合hsdp特性一起使用,否则可能出现精度问题;(反向的权重梯度通信都由hsdp处理)
  • 求反向时,求导函数需使用并行版本parallelize_value_and_grad();
  • 切分中,若涉及partial状态,需通过cell的配置触发重排来消除,或手动消除;
  • 当前仅支持动态图,尚未支持静态图;

测试设计

1)构造一个单卡用例,用例中的网络使用已支持张量并行的算子;
2)使用同一网络,构造一个并行用例;
3)分别跑一个训练step,观察正向loss/反向grad是否能完全对齐;

likedislike
Yyangzhenzhang成员
1月26日 修改了issue 的描述
此处折叠了26条事件消息 查看更多
lijiajia823lijiajia823成员
4月16日 关联了pull request:Update the inference logic and unit tests of some distributed operators