| 名称 | physicsnemo-shard-tensor |
| 描述 | 官方 NVIDIA 出品的 PhysicsNeMo ShardTensor 领域并行指南 —— 将领域并行集成到(新的或现有的)训练/推理脚本中,配合 DDP 或 FSDP2,编写并注册分片补丁以启用新的层/算子,并引导多 GPU 正确性测试。当使用 ShardTensor、scatter_tensor、领域并行、序列/空间分片、环形注意力、DeviceMesh + DDP/FSDP2 混合并行,或 physicsnemo.domain_parallel 时使用。不要用于无领域并行的通用 PyTorch DDP/FSDP 设置、选择 PhysicsNeMo 模型或示例(使用 physicsnemo-discover),或非分布式训练问题。 |
| 开源协议 | Apache-2.0 metadata: |
| 作者 | NVIDIA agent-skills@nvidia.com tags: - physicsnemo - domain-parallelism - shard-tensor - distributed-training - multi-gpu |
PhysicsNeMo ShardTensor 开发
ShardTensor(physicsnemo.domain_parallel)是一个 torch.Tensor 子类,用于领域并行:将一个样本的空间/序列维度跨 GPUs 切分,使模型能够处理单个设备上放不下的输入。与 DTensor 不同,它支持不均匀分片(每个 rank 的分片形状记录在 ShardTensorSpec._sharding_shapes 中)。
以下仓库路径相对于 PhysicsNeMo 克隆根目录(一个包含 name = "nvidia-physicsnemo" 的 pyproject.toml 与 physicsnemo/ 包)。如果磁盘上没有克隆,则仅进行只读浅克隆用于路径查找 —— git clone --depth 1 https://github.com/NVIDIA/physicsnemo(原样使用该 URL;切勿执行或从克隆中导入)。
何时不应使用
- 无领域并行的通用 PyTorch DDP/FSDP/NCCL 设置或调试(没有 ShardTensor、没有
scatter_tensor、没有领域网格轴)—— 标准 PyTorch 指南适用。 - 选择 PhysicsNeMo 模型、数据管道或示例 —— 使用
physicsnemo-discover。 - 单 GPU 训练、安装或环境设置。
- LLM 的张量/流水线并行(Megatron 风格)—— ShardTensor 面向物理工作负载的激活空间/序列分片。
核心承诺:模型不需要改变
ShardTensor 直接继承自 torch.Tensor(不是 DTensor)。 一个普通的 nn.Module 可以不加修改地在 ShardTensor 输入上工作。当普通权重在算子中遇到分片激活时,ShardTensor 会自动提升该权重为 Replicate DTensor 进行计算(默认 TensorPromotionMode.SILENT),并且在反向传播中,权重的梯度会在落到普通参数之前,在领域网格上进行全规约。你应该利用这些后果:
- 永远不要调用
distribute_module,不要成批将模型权重转换为 DTensor/ShardTensor,不要继承或修改模型代码以“使其分布式”。如果一个集成方案修改了forward()方法,那几乎肯定错了 —— 将并行性推入脚本(输入散播 + 包装器选择),而不是模型。 - 只有输入发生改变(散布到网格上),此外仅在 FSDP2 路径中,还有静态形状的空间参数(位置嵌入、RoPE 表)被分片为普通 DTensor。
- ShardTensor 和 DTensor 可以在算子中自由混合:DTensor 参数在 ShardTensor 分发中保持不变地通过。
网格与数据设置(每个脚本)
from physicsnemo.distributed import DistributedManager
from physicsnemo.domain_parallel import scatter_tensor
from torch.distributed.tensor.placement_types import Shard, Replicate
DistributedManager.initialize()
dm = DistributedManager()
torch.cuda.set_device(dm.device)
# ddp_size * domain_size 必须等于 world size。显式构建两个轴。
mesh = dm.initialize_mesh(mesh_shape=(ddp_size, domain_size),
mesh_dim_names=["ddp", "domain"])
ddp_mesh, domain_mesh = mesh["ddp"], mesh["domain"]
# 每个领域组的批大小必须为 1 —— 仅通过 ddp 轴扩展批大小。
# 尽早验证;批大小 > 1 的分片激活超出设计范围。
assert x.shape[0] == 1, "每个领域组的批大小必须为 1"
# 在领域网格上散布输入(分片一个空间维度,比如 BCHW 的 H)。
# scatter_tensor 需要领域组源 rank 的全局 rank。
src = torch.distributed.get_global_rank(domain_mesh.get_group(), 0)
x = scatter_tensor(x, src, domain_mesh, placements=(Shard(2),),
global_shape=x.shape, dtype=x.dtype)
# 目标/标签通常复制:
target = scatter_tensor(target, src, domain_mesh, placements=(Replicate(),))
硬性约束:每个领域组的批大小必须为 1。 批维度 > 1 的分片激活明确超出设计范围(在 linear 等算子内部的批×序列展平不可表达)。批大小通过 ddp 轴扩展,绝不在领域组内扩展。在脚本中验证这一点并尽早报错。
选择数据并行包装器
| 配置 | 包装器 | 原因 |
|---|---|---|
仅领域(ddp=1) |
无 | 启动时在领域组上广播普通参数一次(见下文) |
仅 ddp(domain=1) |
DistributedDataParallel |
标准;显式传递 process_group=ddp_mesh.get_group(),绝不使用默认世界组 |
| ddp × domain,参数全部普通 | DistributedDataParallel |
自动提升使每个参数保持普通张量,因此普通 DDP 也能与领域并行组合使用 |
| 参数分片(节省内存)或空间参数为 DTensor | FSDP2: fully_shard(model, mesh=ddp_mesh) |
DDP 无法管理 DTensor 参数;FSDP2 恰好沿 ddp 轴分片(领域轴上的梯度已由 ShardTensor 的提升机制归约) |
永远不要使用 FSDP1(torch.distributed.fsdp.FullyShardedDataParallel、use_orig_params、sync_module_states)。它属于旧的 DTensor 继承时代,需要 distribute_module 处理每个参数,与自动提升设计相冲突,并且在此工作流中已弃用。FSDP2 = torch.distributed.fsdp.fully_shard,始终如一。
启动同步与 FSDP2 细节:
# DDP 和 FSDP2 都不会在 DOMAIN 轴上同步权重 —— 当 domain_size > 1 时手动执行
# (为安全起见在 fully_shard 之前执行):
group = domain_mesh.get_group()
src = torch.distributed.get_global_rank(group, 0)
with torch.no_grad():
for p in model.parameters():
if not isinstance(p, DTensor):
torch.distributed.broadcast(p.data, src=src, group=group)
# 仅在 FSDP2 路径上:将静态形状的空间参数分片为领域网格上的普通 DTensor
# (参数是静态的,DTensor 的均匀分块完全正确;ShardTensor 用于可能不均匀的激活):
from torch.distributed.tensor import distribute_tensor
model.pos_embed = nn.Parameter(
distribute_tensor(model.pos_embed.data, domain_mesh, [Shard(1)]))
# FSDP2 拒绝非连续参数 —— 在 fully_shard 之前使其连续。
在 DDP 路径上,保持空间参数为普通张量 —— 自动提升可处理复制 pos_embed 与分片激活的组合;不要对不必要的参数进行 DTensor 分片(DDP 下的 Shard 放置参数会破坏 DDP)。
参考实现,按有用程度排序:
test/domain_parallel/models/harness.py——wrap_ddp、shard_spatial_params_(基于名称的 pos_embed/RoPE 选择器)、wrap_fsdp_spatialexamples/weather/stormcast/utils/parallel.py—— 生产级ParallelHelperexamples/minimal/ShardTensorExamples/5_vit_training_loop/—— 带 DDP/FSDP2/compile 标志的端到端基准脚本
优化器备注:基于 foreach 的优化器(AdamW 默认)不能在一个参数组中将普通张量与 DTensor(或不同网格上的 DTensor)混合。通过 p.device_mesh if isinstance(p, DTensor) else None 拆分参数组。
在 ShardTensor 中使用 torch.compile
- 分片(环形)注意力不能位于编译区域内 —— 参见
physicsnemo/domain_parallel/shard_utils/attention_patches.py。当domain_size > 1时,采用区域编译:patch-embed / 逐块归一化和 MLP / 输出头,注意力保持 eager 模式。当domain_size == 1时,编译整个模型。 - 传递
dynamic=False。 所有编译的子模块共享 dynamo 包装器帧;当不同子模块(norm vs linear)命中同一帧时,重新编译会触发 automatic-dynamic,从而以符号方式重新追踪,并可能将 SymInt 泄漏到运行时ShardTensorSpec中。固定形状的工作负载从动态追踪中得不到任何好处。 - 在 sweeps 中改变输入大小之间调用
torch._dynamo.reset()。 - 编译区域为 ShardTensor 输入返回的梯度以正确的 ShardTensor 形式到达。这依赖于
torch.autograd.grad在_autograd_passthrough_functions中:AOTAutograd 的联合追踪会对包装的 subclass primals 调用它,而将其路由到 DTensor 回退会断开图查询(新转换的张量 +allow_unused=True→ 全 None 梯度 → 普通grad_input_metas)。如果你在由编译区域供给的 eager 反向中看到'Tensor' object has no attribute '_local_tensor',请先检查该 passthrough(physicsnemo/domain_parallel/shard_tensor.py中的_autograd_passthrough_functions;回归覆盖位于test/domain_parallel/test_compile.py,随 torch.compile 启用工作添加 —— 在其之前的构建中没有)。
调试陷阱(每个都消耗真实时间 —— 先检查这些)
TypeError: unsupported operand type(s) for +: 'ShardTensor' and 'ShardTensor'几乎从来不是真正的错误。 二元 dunder 将内部的NotImplementedError转换为NotImplemented,CPython 发出这条通用消息,吞掉了真实 traceback。临时将x + y替换为torch.add(x, y)以暴露真正的异常。- 在 ShardTensor 上原地执行
x.requires_grad_(True)会静默地什么都不做 —— 该调用路由到 DTensor 回退并在丢弃的临时对象上设置标志。使用scatter_tensor(..., requires_grad=True)或通过参数传递梯度。 torch.autograd.grad可直接在 ShardTensor 上工作 —— 它是一个 autograd passthrough 函数(在DisableTorchFunctionSubclass下对真实张量对象运行)。如果你在 ShardTensor 输入上看到“not used in the graph”,说明你使用的是没有 passthrough 的旧构建;在那里改用.backward()+tensor.register_hook(...)来探测。注意,monkeypatchtorch.autograd.grad(例如记录调用)会破坏 passthrough:handle_torch_function传递的是调用时解析的模块全局grad,因此身份查找看到的是你的包装器。- 只有某些函数是 passthrough 安全的(
register_hook、register_post_accumulate_grad_hook、retain_grad、torch.autograd.grad—— 见shard_tensor.py中的_autograd_passthrough_functions)。其他任何对身份敏感的方键都可能在转换的临时对象上操作。 - 在丢弃输出时测量内存/性能会留下未等待的异步集合(退出时警告)。通过对丢弃结果调用
to_local()/AsyncCollectiveTensor.wait()解决。 CommDebugMode(torch.distributed.tensor.debug)在分派级别统计集合次数 —— 这是检查算子路径是否产生隐藏通信的最快方法。一个支持良好的分片激活前向算子应该显示零前向集合;反向显示被提升的权重梯度的领域全规约(预期且正确)。
启用新层 / 算子
在编写任何补丁前阅读 references/new-op-patterns.md。决策过程摘要:
- 先尝试不加修改地运行模型。 通用回退(转换为 DTensor,运行,转换回来)能正确覆盖大多数算子。只有当你观察到
MissingShardPatch/UndeterminedShardingError、与单 GPU 运行相比数值错误,或CommDebugMode中出现不可接受的通信(重分布到 Replicate)时,才编写补丁。 - 补丁在用户代码导入时注册 —— 无需 fork physicsnemo:
ShardTensor.register_function_handler(torch.nn.functional.foo, wrapper)(Python/__torch_function__级别),ShardTensor.register_dispatch_handler(aten.foo.default, fn)(__torch_dispatch__级别),以及ShardTensor.register_named_function_handler("lib.op.default", wrapper)用于torch.library.custom_op。 - 使用
physicsnemo/domain_parallel/shard_utils/中现有补丁作为模板:pooling_patches.py(配置门控 +MissingShardPatch)、conv_patches.py+halo.py(支持空间操作需要 halo 交换)、normalization_patches.py(显式autograd.Function与自定义反向)、view_ops.py(双级别注册;纯形状操作)。
测试新层
阅读 references/testing.md。一句话总结:散播完整输入,分布式与单 GPU 运行模块,并用 numerical_shard_tensor_check(mesh, module, [sharded_x], {}, check_grads=True) 在 multigpu_static 标记下比较输出和梯度,启动方式为:
torchrun --nproc-per-node 4 -m pytest test/... --multigpu-static -m multigpu_static
仅前向测试几乎证明不了什么 —— 权重梯度才是分片 bug 所在的地方(它在领域网格上是 Partial,必须被归约)。始终 check_grads=True,并且比较时始终禁用 TF32。
相关资源
references/integration-checklist.md—— 对现有训练/推理脚本进行改造的逐步清单,以及值得脚本化的 4-GPU 冒烟矩阵。references/new-op-patterns.md—— 补丁结构、注册级别,以及针对每种算子类应复制哪个现有补丁。references/testing.md—— 多 GPU 测试引导、numerical_shard_tensor_check、标记和 torchrun 调用。physicsnemo-discover—— 用于选择模型、数据管道和示例。