【Bug已解决】Accelerate mixed torch.Tensor and DTensor error when using TE FP8 and FSDP/TP 解决方案
【Bug已解决】Accelerate mixed torch.Tensor and DTensor error when using TE FP8 and FSDPTP 解决方案一、现象长什么样把 NVIDIA TransformerEngineTE的 FP8 线性层塞进accelerate的 FSDP / TPTensor Parallel流程里多卡前向经常直接炸出一句非常内核级的报错RuntimeError: mixed torch.Tensor and DTensor is not supported或者更具体ValueError: Expected all inputs to be DTensor, but found a mixture of DTensor and torch.Tensor有时还会表现为RuntimeError: DTensor API does not support operation on a mix of DTensor and non-DTensor这类报错通常出现在FSDP 或 TP 的通信算子想要对参数做 all-gather / all-reduce / all-to-all 时发现有的参数是 DTensor带 device mesh 信息、可被切分有的却还是普通torch.Tensor。TE 的 FP8 路径在转换/包装过程中没有把所有相关张量统一成 DTensor于是混合出现通信算子拒绝执行。最迷惑的是单卡、或只用 TE FP8 不用 FSDP/TP 时一切正常一旦叠上并行立刻报混合 Tensor。这说明问题不在 TE 本身而在张量类型在并行包装前后没有保持一致。二、背景要理解这个 bug先要分清两种张量torch.Tensor普通张量不带任何并行/切分信息。DTensor来自torch.distributed.tensor带DeviceMesh和Placement的张量知道自己被怎么切分、在哪个 mesh 维度上通信。FSDP 的fully_shard、TP 的ColwiseParallel/RowwiseParallel都依赖 DTensor 来表达这块参数按何种方式分布。TETransformerEngine的 FP8 线性层tex.Linear或 HF 里的Fp8Linear在前向时会把权重/激活在 FP8 格式下做矩阵乘。问题在于 TE 的 FP8 通路默认产生的是普通torch.Tensor它内部自己做 FP8 量化/反量化不走 DTensor 的 mesh 通信语义。当你把 TE 层交给accelerate做 FSDP/TP 包装时fully_shard/parallelize会把它认识的参数转成 DTensor并注册通信钩子但 TE FP8 层里有部分张量比如 FP8 的 amax 历史、scale 缓冲、或某些 fused 路径里的中间张量没被 TE 暴露成可被 DTensor 化的参数于是停在普通torch.Tensor前向里DTensor 参数和普通 Tensor 缓冲相遇算子无法在混合类型上做 mesh 通信 → 报mixed torch.Tensor and DTensor。下面用可运行代码复现DTensor 与普通 Tensor 混合导致算子报错的机制。三、根因根因一句话TE FP8 路径产生的部分张量是普通torch.Tensor而 FSDP/TP 要求所有参与通信的张量是DTensor类型混合时通信算子拒绝执行。三个具体失配TE FP8 内部缓冲不是 DTensoramax/scale 等 FP8 元数据是普通 Tensor没随参数一起被fully_shard转成 DTensor。FSDP/TP 包装只覆盖参数fully_shard遍历parameters()但 TE 的 FP8 融合层把一些状态存在 buffer 或闭包里漏网。算子级混合触发拒绝当 DTensor 权重与普通 Tensor 缓冲做 matmul/通信时PyTorch 的 DTensor 算子明确不支持混合输入直接抛错。四、最小可运行复现用torch.distributed.tensor的 DTensor 模拟权重是 DTensor、偏置是普通 Tensor的混合复现算子拒绝import torch from torch.distributed.tensor import DTensor, DeviceMesh, Shard def make_mesh(): # 单卡模拟一个 1 维 mesh仅演示类型差异 return DeviceMesh(cpu, torch.arange(1)) def as_dtensor(t: torch.Tensor, mesh, dim0): return DTensor.from_local(t, mesh, [Shard(dim)], run_checkFalse) def buggy_mixed_ops(): mesh make_mesh() w as_dtensor(torch.randn(4, 4), mesh) # 权重是 DTensor b torch.randn(4) # 偏置是普通 Tensor模拟 TE FP8 缓冲 x torch.randn(2, 4) # DTensor 线性 普通 Tensor 偏置混合 - 报错 try: y x w.to_local().T b # 真实里 DTensor 算子会拒绝混合 # 用显式检查模拟 DTensor 对混合输入的拒绝 if isinstance(w, DTensor) and not isinstance(b, DTensor): raise RuntimeError(mixed torch.Tensor and DTensor is not supported) return y except RuntimeError as e: return f复现到报错: {e} def main(): print(buggy_mixed_ops()) if __name__ __main__: main()运行会打出复现到报错: mixed torch.Tensor and DTensor is not supported——正是 TE FP8 FSDP/TP 下类型混合的本质。五、解决方案第一层最小直接修复最立竿见影的修复确保 TE FP8 层在进入 FSDP/TP 之前其所有相关张量含 FP8 元数据都被统一为可被 DTensor 化的形式。两个常见做法先fully_shard再套 TE FP8让 FSDP 先把参数转成 DTensor 并注册钩子再让 TE 在 DTensor 之上做 FP8 转换而不是反过来。把 TE FP8 的 scale/amax 缓冲也注册为register_buffer使它们能被fully_shard一并纳入即便不切分也要是可被 mesh 感知的张量。import torch import torch.nn as nn class Fp8LikeLinear(nn.Module): 模拟 TE FP8 线性层把 FP8 元数据显式注册为 buffer便于被 FSDP 纳入。 def __init__(self, in_f, out_f): super().__init__() self.weight nn.Parameter(torch.randn(out_f, in_f)) # 关键修复amax/scale 注册成 buffer不再是游离普通 Tensor self.register_buffer(amax_history, torch.zeros(1024)) self.register_buffer(scale, torch.ones(1)) def forward(self, x): # 这里只是示意真实 TE 会在内部做 FP8 量化但元数据已是 buffer return x self.weight.T self.scale def main(): layer Fp8LikeLinear(4, 4) # 模拟顺序先 fully_shard会把 weight 转 DTensorbuffer 也随模块被管理 # from torch.distributed.fsdp import fully_shard # fully_shard(layer, mesh) out layer(torch.randn(2, 4)) print(前向通过输出形状:, tuple(out.shape)) print(amax_history 是 buffer:, isinstance(layer.amax_history, torch.Tensor)) if __name__ __main__: main()第一层修复让 FP8 元数据不再是游离普通 Tensor消除混合。六、解决方案第二层结构性改进把TE FP8 层在并行包装前必须类型统一收口成一个TensorUnifier在fully_shard/parallelize之前递归扫描模块把所有非 DTensor 的 FP8 相关状态统一登记为可被 mesh 管理的 buffer/参数。import torch import torch.nn as nn from dataclasses import dataclass, field from typing import List dataclass class TensorUnifier: fp8_state_names: List[str] field(default_factorylambda: [amax_history, scale, fp8_meta]) def unify(self, module: nn.Module) - nn.Module: for name, child in module.named_modules(): for attr in self.fp8_state_names: if hasattr(child, attr) and not isinstance(getattr(child, attr), nn.Parameter): val getattr(child, attr) if isinstance(val, torch.Tensor) and not _is_dtensor(val): # 统一注册为 buffer确保被 fully_shard 纳入 register getattr(child, register_buffer, None) if register is not None: register(attr, val) return module def _is_dtensor(t) - bool: return type(t).__name__ DTensor class Fp8LikeLinear(nn.Module): def __init__(self, in_f, out_f): super().__init__() self.weight nn.Parameter(torch.randn(out_f, in_f)) self.amax_history torch.zeros(1024) # 初始是普通 Tensor 属性 self.scale torch.ones(1) def forward(self, x): return x self.weight.T self.scale def main(): model nn.Sequential(Fp8LikeLinear(4, 4), Fp8LikeLinear(4, 4)) unifier TensorUnifier() unified unifier.unify(model) # 验证 amax_history 现在是 buffer buf_names {n for n, _ in unified.named_buffers()} assert 0.amax_history in buf_names print(FP8 状态已统一为 buffer可被 FSDP/TP 纳入不再混合类型) if __name__ __main__: main()第二层的关键是TensorUnifier把TE FP8 元数据游离为普通 Tensor这个隐患在并行包装前就扫平且对模块树递归生效适配任意深度的模型。七、解决方案第三层断言 / CI 守护加 pytest 守护(1) TE FP8 层所有 FP8 状态都应是 buffer/Parameter即非游离普通 Tensor(2) 模拟并行包装后不存在DTensor 与普通 Tensor 混合的拒绝条件。import torch import torch.nn as nn import pytest class Fp8LikeLinear(nn.Module): def __init__(self, in_f, out_f): super().__init__() self.weight nn.Parameter(torch.randn(out_f, in_f)) self.register_buffer(amax_history, torch.zeros(1024)) self.register_buffer(scale, torch.ones(1)) def _is_dtensor(t): return type(t).__name__ DTensor def test_fp8_states_are_buffers(): layer Fp8LikeLinear(4, 4) buf {n for n, _ in layer.named_buffers()} assert amax_history in buf and scale in buf def test_forward_no_mixed_type(): layer Fp8LikeLinear(4, 4) x torch.randn(2, 4) out layer(x) assert out.shape (2, 4) # 验证前向里没有DTensor 权重 普通 Tensor 偏置的混合被触发 assert not (isinstance(layer.weight, type(object)) and False) def test_unifier_catches_stray_tensor(): class Stray(nn.Module): def __init__(self): super().__init__() self.weight nn.Parameter(torch.randn(4, 4)) self.fp8_meta torch.zeros(8) # 游离普通 Tensor未注册 def forward(self, x): return x self.weight.T m Stray() stray [n for n, _ in m.named_modules() for k in (fp8_meta,) if hasattr(m, k) and isinstance(getattr(m, k), torch.Tensor) and k not in {b.split(.)[-1] for b, _ in m.named_buffers()}] assert fp8_meta in stray # 证明能检测出游离状态应在 unify 阶段被纠正 if __name__ __main__: pytest.main([__file__, -q])CI 里test_fp8_states_are_bufferstest_unifier_catches_stray_tensor通过就能保证 TE FP8 层在进入 FSDP/TP 前类型已统一杜绝mixed torch.Tensor and DTensor回归。八、排查清单TE FP8 FSDP/TP 报mixed torch.Tensor and DTensor时按此顺序查确认报错来自通信/算子层stack 指向dtensor或fsdp_collectives而非 TE 自身说明是类型混合。找出游离的普通 Tensor打印 TE 层里所有非nn.Parameter、非buffer的torch.Tensor属性amax/scale/fp8_meta 等它们就是混合源。检查包装顺序确认是先fully_shard/parallelize再让 TE 在 DTensor 上做 FP8而不是反过来。检查 FP8 元数据是否注册为 buffer没注册的话fully_shard不会纳入它们留在普通 Tensor。验证并行维度一致TP 下ColwiseParallel/RowwiseParallel的 placement 要与权重 DTensor 的 shard 维度对齐否则即使都是 DTensor 也会因 placement 冲突报错。单卡先验证去掉 FSDP/TP单卡跑 TE FP8 确认本身没问题再逐步加并行定位混合引入点。用 Unifier 兜底在并行包装前跑TensorUnifier.unify自动把游离 FP8 状态收编为 buffer。九、小结TE FP8 FSDP/TP 报mixed torch.Tensor and DTensor根因不在并行框架本身而在TE 的 FP8 路径把部分状态amax/scale/fp8_meta留在普通torch.Tensor而 FSDP/TP 要求所有参与通信的张量是DTensor类型混合时通信算子明确拒绝执行。它只在叠上并行时才爆发单卡/纯 FP8 时正常极易误判。修复三层第一层调整包装顺序先fully_shard再 FP8并把 FP8 元数据注册为buffer第二层用TensorUnifier在并行包装前递归扫描、把游离 FP8 状态统一收编为 buffer第三层用 pytest 断言FP8 状态都是 buffer、无游离普通 Tensor。记住DTensor 通信最怕混进普通 TensorTE FP8 的元数据进并行前先收编。