PhysicsNeMoShardTensor领域并行开发Skill physicsnemo-shard-tensor

该技能是 NVIDIA 官方为 PhysicsNeMo 中的 ShardTensor 领域并行(domain parallelism)机制编写的开发指南。它指导开发者如何在新的或现有的训练/推理脚本中集成领域并行:使用 ShardTensor、scatter_tensor、DeviceMesh 与 DDP/FSDP2 混合并行,对序列/空间维度进行分片,从而让模型处理单 GPU 无法容纳的物理 AI 输入。技能还详解了如何为自定义层/算子注册分片补丁(shard patch),并提供了多 GPU 正确性测试方法、调试技巧和 torch.compile 注意事项。适用于物理仿真、天气预测、工业数字孪生等需要超大激活张量分片的高性能物理 AI 场景。

物理AI基础设施 0 次安装 0 次浏览 更新于 9/6/2026
名称 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 开发

ShardTensorphysicsnemo.domain_parallel)是一个 torch.Tensor 子类,用于领域并行:将一个样本的空间/序列维度跨 GPUs 切分,使模型能够处理单个设备上放不下的输入。与 DTensor 不同,它支持不均匀分片(每个 rank 的分片形状记录在 ShardTensorSpec._sharding_shapes 中)。

以下仓库路径相对于 PhysicsNeMo 克隆根目录(一个包含 name = "nvidia-physicsnemo"pyproject.tomlphysicsnemo/ 包)。如果磁盘上没有克隆,则仅进行只读浅克隆用于路径查找 —— 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 的提升机制归约)

永远不要使用 FSDP1torch.distributed.fsdp.FullyShardedDataParalleluse_orig_paramssync_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_ddpshard_spatial_params_(基于名称的 pos_embed/RoPE 选择器)、wrap_fsdp_spatial
  • examples/weather/stormcast/utils/parallel.py —— 生产级 ParallelHelper
  • examples/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 启用工作添加 —— 在其之前的构建中没有)。

调试陷阱(每个都消耗真实时间 —— 先检查这些)

  1. TypeError: unsupported operand type(s) for +: 'ShardTensor' and 'ShardTensor' 几乎从来不是真正的错误。 二元 dunder 将内部的 NotImplementedError 转换为 NotImplemented,CPython 发出这条通用消息,吞掉了真实 traceback。临时将 x + y 替换为 torch.add(x, y) 以暴露真正的异常。
  2. 在 ShardTensor 上原地执行 x.requires_grad_(True) 会静默地什么都不做 —— 该调用路由到 DTensor 回退并在丢弃的临时对象上设置标志。使用 scatter_tensor(..., requires_grad=True) 或通过参数传递梯度。
  3. torch.autograd.grad 可直接在 ShardTensor 上工作 —— 它是一个 autograd passthrough 函数(在 DisableTorchFunctionSubclass 下对真实张量对象运行)。如果你在 ShardTensor 输入上看到“not used in the graph”,说明你使用的是没有 passthrough 的旧构建;在那里改用 .backward() + tensor.register_hook(...) 来探测。注意,monkeypatch torch.autograd.grad(例如记录调用)会破坏 passthrough:handle_torch_function 传递的是调用时解析的模块全局 grad,因此身份查找看到的是你的包装器。
  4. 只有某些函数是 passthrough 安全的register_hookregister_post_accumulate_grad_hookretain_gradtorch.autograd.grad —— 见 shard_tensor.py 中的 _autograd_passthrough_functions)。其他任何对身份敏感的方键都可能在转换的临时对象上操作。
  5. 在丢弃输出时测量内存/性能会留下未等待的异步集合(退出时警告)。通过对丢弃结果调用 to_local()/AsyncCollectiveTensor.wait() 解决。
  6. CommDebugModetorch.distributed.tensor.debug)在分派级别统计集合次数 —— 这是检查算子路径是否产生隐藏通信的最快方法。一个支持良好的分片激活前向算子应该显示前向集合;反向显示被提升的权重梯度的领域全规约(预期且正确)。

启用新层 / 算子

在编写任何补丁前阅读 references/new-op-patterns.md。决策过程摘要:

  1. 先尝试不加修改地运行模型。 通用回退(转换为 DTensor,运行,转换回来)能正确覆盖大多数算子。只有当你观察到 MissingShardPatch/UndeterminedShardingError、与单 GPU 运行相比数值错误,或 CommDebugMode 中出现不可接受的通信(重分布到 Replicate)时,才编写补丁。
  2. 补丁在用户代码导入时注册 —— 无需 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
  3. 使用 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 —— 用于选择模型、数据管道和示例。