在流水想并行场景下,传统的1F1B并行策略虽然进一步优化了流水过程中的Bubble,但是依然有较多Bubble待优化,为了进一步的优化性能,提出了基于dw分离的Zero Bubble并行方案,将计算输入的梯度和计算权重的梯度作了分离, 调整调度方案,增加并行度,提升整体的吞吐量。如图所示,通过分离dx和dw计算,减少流水线并行的空闲时间。
给定一个计算函数,实现一种通用的,业务无关的dx/dw 图分离的方案,增强dw/dx计算的易用性,采用更灵活的调度策略,快速实现dw计算编排。
通过这张正向图为列:
核心流程
1 通过翻转反向计算图,首先获取翻转后每个节点的出度信息,包括{'next_edge': [(grad_node, input_index)]}, 2. input_index即获取父节点的边的信息,用于路径选择。 3. 以输入节点为截止条件,根节点为首节点,BFS遍历反向计算图,得到dx的计算子图。 获取sub_graph的set集合,用于后续计算最邻近公共节点。
通过BFS算法,以每一个权重为起点,找到每一个权重的最邻近公共节点,以及连边信息,用于路径剪枝。
使用并查集思想,合并公共父节点路径,保证最小子树无交叉
1.获取到所有的 子图后,需要进行路径剪枝,避免子图多个根节点之间存在父子关系导致梯度重复累加,具体实现上,根据合并后的 param_group {"w":{w1}, "immediate": {a1, a2}, "edge_index": {a1:{0}, a2:{1}}} 以上述有重叠场景为例,根据param_group 的信息,我们需要把a2中间节点的第0条边置空,使用grad_node._set_next_edge(index, None)实现剪枝。分组计算和不分组差异 主要在显存可以及时释放,避免出现显存峰值上升。
forward_and_gradfn(fn, *inputs, weights=None, has_aux=False, grad_position=0, **kwargs)
forward_out, grad_fn
GradFunction
compute_input_grad(sens=None) -> dx
compute_weight_grad(keep_graph=False) -> dw
__call__(sens=None, keep_graph=False)
compute_weight_grad
compute_input_grad
在网络加了@jit后,能够按照正常的dw/dx分离算法执行,并且在性能上有提升,如果网络不支持直接加@jit,由动态图兜底dw/dx分离。
keep_graph=True
mint.split
_GradientEdge
当前先支持PyNative
需求来源
在流水想并行场景下,传统的1F1B并行策略虽然进一步优化了流水过程中的Bubble,但是依然有较多Bubble待优化,为了进一步的优化性能,提出了基于dw分离的Zero Bubble并行方案,将计算输入的梯度和计算权重的梯度作了分离, 调整调度方案,增加并行度,提升整体的吞吐量。如图所示,通过分离dx和dw计算,减少流水线并行的空闲时间。
目标
给定一个计算函数,实现一种通用的,业务无关的dx/dw 图分离的方案,增强dw/dx计算的易用性,采用更灵活的调度策略,快速实现dw计算编排。

设计思路
通过这张正向图为列:

算法实现
核心流程
通过上述流程获取dx,dw的计算子图如下:
dx计算图
dw计算图
dx路径闭包计算
1 通过翻转反向计算图,首先获取翻转后每个节点的出度信息,包括{'next_edge': [(grad_node, input_index)]},
2. input_index即获取父节点的边的信息,用于路径选择。
3. 以输入节点为截止条件,根节点为首节点,BFS遍历反向计算图,得到dx的计算子图。
获取sub_graph的set集合,用于后续计算最邻近公共节点。
计算最最邻近公共子节点
通过BFS算法,以每一个权重为起点,找到每一个权重的最邻近公共节点,以及连边信息,用于路径剪枝。
相同公共节点合并
使用并查集思想,合并公共父节点路径,保证最小子树无交叉
最近邻公共节点注册prehook,计算dx子图
路径剪枝,分组计算dw
1.获取到所有的 子图后,需要进行路径剪枝,避免子图多个根节点之间存在父子关系导致梯度重复累加,具体实现上,根据合并后的
param_group {"w":{w1}, "immediate": {a1, a2}, "edge_index": {a1:{0}, a2:{1}}} 以上述有重叠场景为例,根据param_group
的信息,我们需要把a2中间节点的第0条边置空,使用grad_node._set_next_edge(index, None)实现剪枝。分组计算和不分组差异
主要在显存可以及时释放,避免出现显存峰值上升。
涉及到的对外API
接口设计
forward_and_gradfn(fn, *inputs, weights=None, has_aux=False, grad_position=0, **kwargs)forward_out, grad_fn6.2
GradFunctioncompute_input_grad(sens=None) -> dxcompute_weight_grad(keep_graph=False) -> dw__call__(sens=None, keep_graph=False):按配置返回 dx/dw 或组合结果约束
compute_weight_grad必须在compute_input_grad之后调用(依赖已捕获的 intermediates grads)。与其他模块的相关性描述
动静统一方案
在网络加了@jit后,能够按照正常的dw/dx分离算法执行,并且在性能上有提升,如果网络不支持直接加@jit,由动态图兜底dw/dx分离。
测试设计与测试计划
测试用例设计
keep_graph=True:dw 可重复调用且一致mint.split,按 slot 处理_GradientEdge其他信息
当前先支持PyNative