PyTorch ColwiseParallel / RowwiseParallel / SequenceParallel 在模块 I/O 边界调用 DTensor.redistribute(async_op=True)。collective 提交后不立即 wait,而是返回 torch.distributed._functional_collectives.AsyncCollectiveTensor(ACT),在下游算子(如 Linear)第一次读取 local tensor 时才 wait_tensor(),从而实现通信与计算重叠。
ColwiseParallel
RowwiseParallel
SequenceParallel
DTensor.redistribute(async_op=True)
wait
torch.distributed._functional_collectives.AsyncCollectiveTensor
Linear
wait_tensor()
Hyper parallelize_module 的 Col/Row/Seq 当前走 DTensor.redistribute() 全同步路径:TensorRedistribution → differentiable_all_gather_concat / differentiable_all_reduce 等均在返回前完成通信。功能正确,但无 TP 边界重叠。
parallelize_module
DTensor.redistribute()
TensorRedistribution
differentiable_all_gather_concat
differentiable_all_reduce
本 Issue 目标:在 Torch 后端,基于 PyTorch AsyncCollectiveTensor + wait_tensor(),为 Col/Row/Seq 的 redistribute 增加 async_op 能力。
AsyncCollectiveTensor
async_op
关联: changzherui1/hyper-parallel#6 TP 专项 · mindspore/hyper-parallel#269 loss_parallel
参考实现(PyTorch 侧):
async_op=True
torch/distributed/tensor/parallel/style.py
Redistribute
torch/distributed/tensor/_redistribute.py
wait_tensor
torch/distributed/_functional_collectives.py
Hyper 侧现状:
hyper_parallel/core/tensor_parallel/style.py
redistribute()
hyper_parallel/core/dtensor/tensor_redistribution.py
hyper_parallel/platform/torch/platform.py
all_gather_single(..., async_op=True)
differentiable_all_to_all_single_async
本 Issue 范围: Torch 后端 + AsyncCollectiveTensor。MindSpore 走自有 AsyncCollectiveTensor(platform/mindspore/platform.py),另开子任务,不在本 Issue 首版交付。
platform/mindspore/platform.py
redistribute(async_op=True) → funcol.all_gather_single / all_reduce / ... → _wrap_tensor_autograd(elem) → AsyncCollectiveTensor → 返回 DTensor(_local_tensor=ACT, ...) 下游 Linear 读 input: → ACT.__torch_dispatch__ → trigger_wait() → wait_tensor(elem) → 与已提交的 collective 重叠 async_op=False(默认): → redistribute_local_tensor 末尾 new_local_tensor.wait()
关键约束:
__torch_dispatch__
to_local()
_local_tensor
nn.Linear
结论: TP 异步不需要给 Hyper DTensorBase 实现全套 __torch_dispatch__;复用 PyTorch AsyncCollectiveTensor.__torch_dispatch__ 即可。
DTensorBase
AsyncCollectiveTensor.__torch_dispatch__
__torch_function__
_OP_DISPATCHER
platform/torch/dtensor.py
DTensor.__torch_dispatch__
redistribute
trigger_wait()
__ms_dispatch__
platform/mindspore/...
路径 A(推荐,本 Issue): redistribute(async_op=True) → DTensor._local_tensor = ACT use_local_output=True: to_local() → ACT → nn.Linear → ACT.__torch_dispatch__ wait 路径 B(不推荐): 给 Hyper DTensorBase 加 __torch_dispatch__ → 与 OpDispatcher 架构冲突 路径 C(需验证): use_local_output=False → DTensor op → _OP_DISPATCHER._unwrap_value → to_local() 取 ACT → 已注册 op 调用时 ACT dispatch 应仍生效;full-gather 回退路径需避免 eager wait
与 #269 关系: loss_parallel CE 走融合 kernel + _OP_DISPATCHER,不依赖本 Issue 的 redistribute 重叠;两者可并行推进。
all_gather
DTensor.to_local()
parallelize_module → Col/Row/Seq._prepare_input_fn / _prepare_output_fn → DTensor.redistribute(..., async_op=True) # 新增参数,默认 False → TensorRedistribution.redistribution(..., async_op) → all_concat / all_reduce / all_to_all(platform 层) → async_op=True: 返回 ACT,不 wait → async_op=False: work.wait()(保持现状)
原则: async_op=False 为默认,不改变现有行为;仅 TP style 显式传 True。
async_op=False
True
文件: hyper_parallel/platform/torch/platform.py
新增(或扩展):
def differentiable_all_gather_concat_async(data, group, concat_size, concat_dim, rank_list=None): # all_gather_into_tensor(..., async_op=True) # 用 funcol 路径返回 ACT(参考 differentiable_all_to_all_single_async) ... def differentiable_all_reduce_async(data, op, group): # all_reduce(..., async_op=True) → ACT ...
复用现有:
wait_async_tensor()
_wrap_tensor_autograd
注意: differentiable_all_gather_concat 当前用 list(dist.all_gather);异步版建议改为 all_gather_into_tensor 单 buffer,与 FSDP 路径一致。
list(dist.all_gather)
all_gather_into_tensor
文件: hyper_parallel/core/dtensor/tensor_redistribution.py
redistribution(self, input_x, to_layout, *, async_op=False)
_construct_all_concat
_construct_all_concat_new
all_reduce
all_to_all
与 PyTorch 对齐的逻辑(单步 collective 后):
if not async_op and isinstance(result, AsyncCollectiveTensor): result = result.wait()
DTensor.redistribute(async_op=False)
文件: hyper_parallel/core/dtensor/dtensor.py
hyper_parallel/core/dtensor/dtensor.py
def redistribute(self, device_mesh, placements, *, async_op=False) -> DTensor: ... out = _tensor_redistribution.redistribution(self, dst_layout, async_op=async_op)
方案 A(推荐,与 PyTorch 一致): 新增 RedistributeFunction(torch.autograd.Function),forward/backward 保存 async_op,backward 调用反向 redistribution 时同样 async_op=True。
RedistributeFunction(torch.autograd.Function)
方案 B(最小 PoC): 首版仅 forward async,backward 同步——有重叠收益但 backward 无重叠。
首版建议 方案 A,至少覆盖 Col/Row 主路径。
文件: hyper_parallel/core/tensor_parallel/style.py
与 PyTorch 相同位置传参:
_prepare_input_fn
_prepare_output_fn
PrepareModuleInput/Output
to_local
文件: hyper_parallel/core/shard/_op_dispatch.py、platform/torch/dtensor.py
hyper_parallel/core/shard/_op_dispatch.py
Hyper OpDispatcher._unwrap_value 对 DTensor 调用 to_local(),不会对 ACT 提前 wait()——有利于 ACT 透传。需验证:
OpDispatcher._unwrap_value
wait()
use_local_output=True
use_local_output=False
SkipDTensorDispatch
若 op dispatch 对 ACT 提前 wait,需在 dispatch 入口 透传 ACT 而非 eager wait(仅当实测重叠失效时再改)。
isinstance(local, AsyncCollectiveTensor)
tensor_parallel
测试文件建议:
tests/torch/tensor_parallel/test_tp_redistribute_async.py
examples/torch/llama3/tensor_parallel_example.py
--async-redistribute
redistribute(async_op)
不在本 Issue:
torch.compile
PrepareModuleInput
# style.py — ColwiseParallel input_tensor.redistribute(placements=desired_input_layouts, async_op=True) outputs.redistribute(placements=output_layouts, async_op=True) # _redistribute.py if not async_op and isinstance(new_local_tensor, funcol.AsyncCollectiveTensor): new_local_tensor = new_local_tensor.wait() # _functional_collectives.py — AsyncCollectiveTensor.__torch_dispatch__ # 非 view op → trigger_wait()
关联:changzherui1/hyper-parallel#6 · mindspore/hyper-parallel#269
一、背景
PyTorch
ColwiseParallel/RowwiseParallel/SequenceParallel在模块 I/O 边界调用DTensor.redistribute(async_op=True)。collective 提交后不立即wait,而是返回torch.distributed._functional_collectives.AsyncCollectiveTensor(ACT),在下游算子(如Linear)第一次读取 local tensor 时才wait_tensor(),从而实现通信与计算重叠。Hyper
parallelize_module的 Col/Row/Seq 当前走DTensor.redistribute()全同步路径:TensorRedistribution→differentiable_all_gather_concat/differentiable_all_reduce等均在返回前完成通信。功能正确,但无 TP 边界重叠。本 Issue 目标:在 Torch 后端,基于 PyTorch
AsyncCollectiveTensor+wait_tensor(),为 Col/Row/Seq 的 redistribute 增加async_op能力。关联: changzherui1/hyper-parallel#6 TP 专项 · mindspore/hyper-parallel#269 loss_parallel
参考实现(PyTorch 侧):
async_op=Truetorch/distributed/tensor/parallel/style.pyRedistributeautogradtorch/distributed/tensor/_redistribute.pywait_tensortorch/distributed/_functional_collectives.pyHyper 侧现状:
hyper_parallel/core/tensor_parallel/style.pyredistribute()无async_ophyper_parallel/core/dtensor/tensor_redistribution.pyhyper_parallel/platform/torch/platform.pyall_gather_single(..., async_op=True)已有;differentiable_all_gather_concat同步differentiable_all_to_all_single_async、FSDP async AG本 Issue 范围: Torch 后端 +
AsyncCollectiveTensor。MindSpore 走自有AsyncCollectiveTensor(platform/mindspore/platform.py),另开子任务,不在本 Issue 首版交付。二、PyTorch 异步机制摘要
关键约束:
__torch_dispatch__延迟 wait;view 类 op 可能不 wait(PyTorch 对 view 有特殊处理)。async_op,否则只有 forward 重叠。to_local()若直接返回_local_tensor而不经 autograd,ACT 会原样传给nn.Linear——这正是期望行为。二点五、Hyper 与
__torch_dispatch__的关系__torch_function__→_OP_DISPATCHER(platform/torch/dtensor.py)DTensor.__torch_dispatch__redistribute→ ACTAsyncCollectiveTensor__torch_dispatch__→trigger_wait()__ms_dispatch__(platform/mindspore/...)与 #269 关系: loss_parallel CE 走融合 kernel +
_OP_DISPATCHER,不依赖本 Issue 的 redistribute 重叠;两者可并行推进。三、Hyper 当前差异
DTensor.redistribute()无async_op参数TensorRedistribution全部同步 waitdifferentiable_all_gather_concat用同步all_gatherDTensor.to_local()直接返回_local_tensorasync_op=True四、实现方案
4.1 总体架构
原则:
async_op=False为默认,不改变现有行为;仅 TP style 显式传True。4.2 分步任务
Step 1 — Platform:异步 all_gather / all_reduce 包装 ACT
文件:
hyper_parallel/platform/torch/platform.py新增(或扩展):
def differentiable_all_gather_concat_async(data, group, concat_size, concat_dim, rank_list=None): # all_gather_into_tensor(..., async_op=True) # 用 funcol 路径返回 ACT(参考 differentiable_all_to_all_single_async) ... def differentiable_all_reduce_async(data, op, group): # all_reduce(..., async_op=True) → ACT ...复用现有:
wait_async_tensor()→wait_tensor()_wrap_tensor_autograd/AsyncCollectiveTensor(PyTorch 内置)注意:
differentiable_all_gather_concat当前用list(dist.all_gather);异步版建议改为all_gather_into_tensor单 buffer,与 FSDP 路径一致。Step 2 —
TensorRedistribution传递async_op文件:
hyper_parallel/core/dtensor/tensor_redistribution.pyredistribution(self, input_x, to_layout, *, async_op=False)_construct_all_concat/_construct_all_concat_new:根据async_op分支同步/异步all_reduce(partial → replicate)路径同理all_to_all:可复用differentiable_all_to_all_single_async或新增 async 变体与 PyTorch 对齐的逻辑(单步 collective 后):
if not async_op and isinstance(result, AsyncCollectiveTensor): result = result.wait()Step 3 —
DTensor.redistribute(async_op=False)文件:
hyper_parallel/core/dtensor/dtensor.pydef redistribute(self, device_mesh, placements, *, async_op=False) -> DTensor: ... out = _tensor_redistribution.redistribution(self, dst_layout, async_op=async_op)Step 4 — Autograd:backward 支持 async
方案 A(推荐,与 PyTorch 一致): 新增
RedistributeFunction(torch.autograd.Function),forward/backward 保存async_op,backward 调用反向 redistribution 时同样async_op=True。方案 B(最小 PoC): 首版仅 forward async,backward 同步——有重叠收益但 backward 无重叠。
首版建议 方案 A,至少覆盖 Col/Row 主路径。
Step 5 — TP style 开启
async_op=True文件:
hyper_parallel/core/tensor_parallel/style.py与 PyTorch 相同位置传参:
ColwiseParallel_prepare_input_fn、_prepare_output_fn的redistributeRowwiseParallelSequenceParallel_prepare_input_fn的redistributePrepareModuleInput/OutputStep 6 — Op dispatch /
to_local验证文件:
hyper_parallel/core/shard/_op_dispatch.py、platform/torch/dtensor.pyHyper
OpDispatcher._unwrap_value对 DTensor 调用to_local(),不会对 ACT 提前wait()——有利于 ACT 透传。需验证:use_local_output=True:to_local()→ ACT →nn.Linear→ ACT__torch_dispatch__wait ✅use_local_output=False:DTensor 包裹 ACT →_OP_DISPATCHER→ unwrap 后 local 仍为 ACT → 已注册 op 应触发 ACT waitSkipDTensorDispatchbackward:plain tensor 路径仍靠 ACT 自身 dispatch若 op dispatch 对 ACT 提前
wait,需在 dispatch 入口 透传 ACT 而非 eager wait(仅当实测重叠失效时再改)。五、TP 边界 collective 映射
六、测试与验收
isinstance(local, AsyncCollectiveTensor)在 hook 后、Linear 前为 Truetensor_parallelST 全过(默认async_op=False)测试文件建议:
tests/torch/tensor_parallel/test_tp_redistribute_async.py(新建)examples/torch/llama3/tensor_parallel_example.py可选--async-redistribute七、里程碑
redistribute(async_op)forward only + Colwise八、风险与不在范围
wait_tensor与 stream 行为不在本 Issue:
AsyncCollectiveTensor路径(另开)torch.compile/ AsyncTP(Inductor micro-pipeline)PrepareModuleInputasync(PyTorch 亦未默认开启)附录:PyTorch 参考代码位置
# style.py — ColwiseParallel input_tensor.redistribute(placements=desired_input_layouts, async_op=True) outputs.redistribute(placements=output_layouts, async_op=True) # _redistribute.py if not async_op and isinstance(new_local_tensor, funcol.AsyncCollectiveTensor): new_local_tensor = new_local_tensor.wait() # _functional_collectives.py — AsyncCollectiveTensor.__torch_dispatch__ # 非 view op → trigger_wait()关联:changzherui1/hyper-parallel#6 · mindspore/hyper-parallel#269