PyTorch LocalTensor 实战教程:单进程 SPMD 分布式调试
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
本教程围绕 PyTorch 分布式调试工具LocalTensor展开:它是一个在单进程内模拟 SPMD(Single Program, Multiple Data)分布式计算的torch.Tensor子类,让你无需拉起多进程或多张 GPU,就能对all_reduce、broadcast等集合通信操作以及 DTensor 张量并行代码进行本地调试与断言验证。读完本文,你将掌握LocalTensor/LocalTensorMode/LocalIntNode的完整用法、六类典型调试场景的可复制代码,以及非均匀分片与多维 mesh 混合并行的测试技巧。
1. 什么是 LocalTensor
LocalTensor是torch.Tensor的子类,它在单个进程内部维护一个「rank ID → 本地分片张量」的映射,从而在一个进程里模拟出分布式 SPMD 计算。源码中它的定义位于 torch/distributed/_local_tensor/__init__.py(class LocalTensor(torch.Tensor)),每个 rank 对应的分片存放在_local_tensors字典中,可以通过tensor._local_tensors[rank]直接查看任意 rank 的状态。
核心优势:
- 无需多进程环境:在单进程 CPU/GPU 上即可测试分布式算法;
- 调试迭代更快:不必反复拉起多进程;
- 全量可见性:直接检查每个 rank 的张量状态;
- CI 友好:可以放入单进程 CI 流水线中跑分布式测试;
- DTensor 集成:无缝在本地测试 DTensor 分布式张量代码。
注意:
LocalTensor仅用于调试与测试,不适合生产环境。在本地模拟多个 rank 会带来显著开销(源码中每个算子都要按 rank 逐个分发执行,见__torch_dispatch__,torch/distributed/_local_tensor/__init__.py)。
安装与环境准备
LocalTensor属于 PyTorch 分布式包的一部分,安装 PyTorch 后即可使用,无需额外依赖:
from torch.distributed._local_tensor import ( LocalTensor, LocalTensorMode, LocalIntNode, )需要注意:即使只做本地测试,涉及集合通信(dist.all_reduce等)时仍需要初始化一个进程组,各示例统一使用"fake"后端:
import torch.distributed as dist dist.init_process_group("fake", rank=0, world_size=3) pg = dist.distributed_c10d._get_default_group() # ... 使用 pg 执行集合通信 ... dist.destroy_process_group()以下所有示例代码均直接来自仓库中可执行、可测试的源码文件(各文件均可通过python <file>直接运行验证),测试套件也会调用这些相同函数保证文档正确性。
2. 示例 1:创建与基础运算
完整源码见 example_01_basic_operations.py。
从各 rank 张量创建 LocalTensor:
def create_local_tensor(): """Create a LocalTensor from per-rank tensors.""" rank_0_tensor = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) rank_1_tensor = torch.tensor([[5.0, 6.0], [7.0, 8.0]]) local_tensor = LocalTensor({0: rank_0_tensor, 1: rank_1_tensor}) expected = (torch.Size([2, 2]), frozenset({0, 1}), rank_0_tensor, rank_1_tensor) return local_tensor, expected构造后local_tensor.shape为torch.Size([2, 2]),local_tensor._ranks为frozenset({0, 1})——各 rank 分片形状一致时,shape返回普通torch.Size。
逐 rank 的算术运算:
def arithmetic_operations(): """Demonstrate arithmetic on LocalTensor.""" input_0 = torch.tensor([1.0, 2.0, 3.0]) input_1 = torch.tensor([4.0, 5.0, 6.0]) lt = LocalTensor({0: input_0, 1: input_1}) doubled = lt * 2 # 对每个 rank 的分片执行同一算子(SPMD 语义) added = lt + 10 expected = (input_0 * 2, input_1 * 2, input_0 + 10) return (doubled, added), expectedlt * 2的结果仍是LocalTensor,其_local_tensors[0]等于input_0 * 2、_local_tensors[1]等于input_1 * 2。
所有分片一致时提取单个张量(reconcile):
def reconcile_identical_shards(): """Extract a single tensor when all shards are identical.""" value = torch.tensor([1.0, 2.0, 3.0]) lt = LocalTensor({0: value.clone(), 1: value.clone(), 2: value.clone()}) result = lt.reconcile() # 返回普通 torch.Tensor return result, valuereconcile()的实现位于 torch/distributed/_local_tensor/__init__.py,当所有 rank 分片数值相同时(例如 all-reduce 之后)将其「收敛」为一个普通张量;分片不一致时行为按内部判定(可能返回None),它非常适合用于断言。
使用 LocalTensorMode 自动创建 LocalTensor:
def use_local_tensor_mode(world_size: int = 4): """Use LocalTensorMode to auto-create LocalTensors.""" with LocalTensorMode(world_size): x = torch.ones(2, 3) # 工厂函数自动产出 LocalTensor is_local = isinstance(x, LocalTensor) num_ranks = len(x._ranks) return (is_local, num_ranks), (True, world_size)LocalTensorMode是一个TorchDispatchMode(定义见 torch/distributed/_local_tensor/__init__.py),进入上下文后torch.ones、torch.tensor等工厂函数会直接产出带world_size个 rank 的LocalTensor。
直接访问各 rank 分片调试:
def access_individual_shards(): """Access shards for debugging.""" input_0 = torch.tensor([1.0, 2.0]) input_1 = torch.tensor([3.0, 4.0]) lt = LocalTensor({0: input_0, 1: input_1, 2: torch.tensor([5.0, 6.0])}) shard_0 = lt._local_tensors[0] # 字典方式 shard_1 = lt._local_tensor_1 # 属性方式(_local_tensor_<rank>) return (shard_0, shard_1), (input_0, input_1)3. 示例 2:模拟集合通信操作
完整源码见 example_02_collective_operations.py。该文件的__main__部分展示了标准流程:先dist.init_process_group("fake", rank=0, world_size=3),取默认进程组pg,跑完后destroy_process_group()。
All-reduce(SUM):
def all_reduce_sum(process_group): """Simulate all_reduce with SUM across ranks.""" tensors = { 0: torch.tensor([[1.0, 2.0], [3.0, 4.0]]), 1: torch.tensor([[5.0, 6.0], [7.0, 8.0]]), 2: torch.tensor([[9.0, 10.0], [11.0, 12.0]]), } expected = sum(tensors.values()) with LocalTensorMode(frozenset(tensors.keys())): lt = LocalTensor({k: v.clone() for k, v in tensors.items()}) dist.all_reduce(lt, op=dist.ReduceOp.SUM, group=process_group) result = lt.reconcile() # 三个 rank 的值现在相同,可收敛为单张量 return result, expected从指定 rank 广播:
def broadcast_from_rank(process_group, src_rank: int = 0): """Simulate broadcast from a source rank.""" tensors = { 0: torch.tensor([10.0, 20.0, 30.0]), 1: torch.tensor([40.0, 50.0, 60.0]), 2: torch.tensor([70.0, 80.0, 90.0]), } expected = tensors[src_rank].clone() with LocalTensorMode(frozenset(tensors.keys())): lt = LocalTensor({k: v.clone() for k, v in tensors.items()}) dist.broadcast(lt, src=src_rank, group=process_group) result = lt.reconcile() return result, expectedAll-gather 收集所有 rank 的张量:
def all_gather_tensors(process_group): """Simulate all_gather to collect tensors from all ranks.""" tensors = { 0: torch.tensor([[1.0, 2.0]]), 1: torch.tensor([[3.0, 4.0]]), 2: torch.tensor([[5.0, 6.0]]), } num_ranks = len(tensors) expected = [tensors[i].clone() for i in range(num_ranks)] with LocalTensorMode(frozenset(tensors.keys())): lt = LocalTensor(tensors) output_list = [torch.zeros_like(lt) for _ in range(num_ranks)] dist.all_gather(output_list, lt, group=process_group) results = [out.reconcile() for out in output_list] return results, expected同一文件中还有一个reduce_scatter_tensors()示例(先归约再按行切分回各 rank),展示了dist.reduce_scatter_single(lt_output, lt_input, group=process_group)的模拟用法,可直接参考 example_02_collective_operations.py。
4. 示例 3:与 DTensor 集成
完整源码见 example_03_dtensor_integration.py。LocalTensor与 DTensor 配合后可在本地测试分布式张量并行。
分发张量并验证重建:
def distribute_and_verify(world_size: int = 4): """Distribute a tensor and verify reconstruction.""" with LocalTensorMode(world_size): mesh = init_device_mesh("cpu", (world_size,)) tensor = torch.arange(16).reshape(4, 4).float() dt_sharded = distribute_tensor(tensor, mesh, [Shard(0)]) dt_replicated = distribute_tensor(tensor, mesh, [Replicate()]) sharded_actual = dt_sharded.full_tensor().reconcile() replicated_actual = dt_replicated.to_local().reconcile() return (sharded_actual, replicated_actual), (tensor, tensor)distribute_tensor(tensor, mesh, [Shard(0)])沿第 0 维切分到 4 个 rank;full_tensor()聚合回全局张量后,用.reconcile()取回普通张量与原张量比对。Replicate()则让每个 rank 持有完整副本,此时to_local()在每个 rank 上都等于原张量。
分布式矩阵乘法:
def dtensor_matmul(world_size: int = 4): """Perform matrix multiplication with DTensors.""" with LocalTensorMode(world_size): mesh = init_device_mesh("cpu", (world_size,)) a = torch.randn(8, 4) b = torch.randn(4, 6) da = distribute_tensor(a, mesh, [Shard(0)]) # 行并行:输入按行切 db = distribute_tensor(b, mesh, [Replicate()]) # 权重全量复制 dc = da @ db # DTensor 自动推导输出布局 expected = a @ b actual = dc.full_tensor().reconcile() return actual, expected模拟分布式 Linear 层前向:
def dtensor_linear_layer(world_size: int = 4): """Simulate a distributed linear layer forward pass.""" batch_size, in_features, out_features = 16, 8, 4 with LocalTensorMode(world_size): mesh = init_device_mesh("cpu", (world_size,)) x = torch.randn(batch_size, in_features) w = torch.randn(in_features, out_features) b = torch.randn(out_features) dx = distribute_tensor(x, mesh, [Shard(0)]) dw = distribute_tensor(w, mesh, [Replicate()]) db = distribute_tensor(b, mesh, [Replicate()]) dy = torch.relu(dx @ dw + db) expected = torch.relu(x @ w + b) actual = dy.full_tensor().reconcile() return actual, expected这套模式正是数据并行训练的本地验证范式:激活按 batch 维Shard(0),参数Replicate(),前向结果与单进程参考实现逐元素比对(测试中用torch.allclose(actual, expected, atol=1e-5))。
5. 示例 4:处理非均匀分片
真实分布式系统中各 rank 的数据量常常不相等(例如总行数不能被 world_size 整除)。LocalTensor通过SymInt形状与LocalIntNode处理这种情况。完整源码见 example_04_uneven_sharding.py。
创建各 rank 尺寸不同的 LocalTensor:
def create_uneven_shards(): """Create LocalTensor with different sizes per rank.""" tensors = { 0: torch.tensor([[1.0, 2.0, 3.0, 4.0]]), # 1 行 1: torch.tensor([[5.0, 6.0, 7.0, 8.0], [9.0, 10.0, 11.0, 12.0]]), # 2 行 2: torch.tensor([[13.0, 14.0, 15.0, 16.0]]), # 1 行 } lt = LocalTensor(tensors) is_symint = isinstance(lt.shape[0], torch.SymInt) expected_shapes = {rank: t.shape for rank, t in tensors.items()} return (lt, is_symint), expected_shapes各 rank 第 0 维不一致时,lt.shape[0]不再是普通int,而是一个torch.SymInt(每个 rank 持有各自的符号值),后续算子会基于符号形状做形状推导。
LocalIntNode 逐 rank 整型运算:
def local_int_node_arithmetic(): """LocalIntNode for per-rank integer values.""" values_a = {0: 10, 1: 20, 2: 30} values_b = {0: 1, 1: 2, 2: 3} local_a = LocalIntNode(values_a) local_b = LocalIntNode(values_b) result_add = local_a.add(local_b) result_mul = local_a.mul(local_b) expected_add = {k: values_a[k] + values_b[k] for k in values_a} expected_mul = {k: values_a[k] * values_b[k] for k in values_a} return ( (dict(result_add._local_ints), dict(result_mul._local_ints)), (expected_add, expected_mul), )LocalIntNode定义于 torch/distributed/_local_tensor/__init__.py,内部以_local_ints字典保存各 rank 的整型值,支持add/sub/mul/floordiv/mod、比较运算(eq/ge/lt等,结果可为bool | SymBool)以及sym_max/sym_min/sym_sum等符号运算;当各 rank 值相同时会退化/兼容为普通int(ConstantIntNode)。它是 DTensor 做非均匀切分形状计算的底层支撑。
DTensor 处理不能整除的维度:
def dtensor_uneven_sharding(world_size: int = 3): """DTensor with unevenly divisible tensor dimension.""" total_rows = 10 # 10 行切给 3 个 rank:4 + 3 + 3 with LocalTensorMode(world_size): mesh = init_device_mesh("cpu", (world_size,)) tensor = torch.arange(total_rows * 4).reshape(total_rows, 4).float() dt = distribute_tensor(tensor, mesh, [Shard(0)]) local = dt.to_local() rows_per_rank = { rank: local._local_tensors[rank].shape[0] for rank in range(world_size) } reconstructed = dt.full_tensor().reconcile() matches = torch.equal(reconstructed, tensor) return (rows_per_rank, matches), total_rows测试断言sum(rows_per_rank.values()) == 10且重建后与原张量完全相等——这正是非均匀分片下「切分 + 聚合无损」的关键验证点。
6. 示例 5:Rank 专属计算(非 SPMD 行为)
有时你需要对不同 rank 执行不同的逻辑(而非 SPMD 的同一操作)。完整源码见 example_05_rank_specific.py。
用 rank_map() 创建逐 rank 的值:
def use_rank_map(world_size: int = 4): """Create LocalTensors with per-rank values using rank_map.""" with LocalTensorMode(world_size): lt = rank_map(lambda rank: torch.full((2, 3), float(rank))) values = { rank: lt._local_tensors[rank][0, 0].item() for rank in range(world_size) } expected = {rank: float(rank) for rank in range(world_size)} return values, expected用 tensor_map() 对每个 rank 的分片做不同变换:
def use_tensor_map(world_size: int = 4): """Transform each shard differently using tensor_map.""" with LocalTensorMode(world_size): lt = rank_map(lambda rank: torch.ones(2, 2) * (rank + 1)) def scale_by_rank(rank: int, tensor: torch.Tensor) -> torch.Tensor: return tensor * (rank + 1) scaled = tensor_map(lt, scale_by_rank) values = { rank: scaled._local_tensors[rank][0, 0].item() for rank in range(world_size) } # (rank + 1) * (rank + 1) = (rank + 1)^2 expected = {rank: float((rank + 1) ** 2) for rank in range(world_size)} return values, expectedrank_map(torch/distributed/_local_tensor/__init__.py)接收一个rank -> Tensor回调并逐 rank 执行;tensor_map(同文件 L1877)接收(rank, shard) -> Tensor回调,对已有 LocalTensor 的各分片分别变换。
临时退出 LocalTensorMode:
def disable_mode_temporarily(world_size: int = 4): """Temporarily exit LocalTensorMode for regular tensor ops.""" with LocalTensorMode(world_size) as mode: lt = torch.ones(2, 2) inside_type = type(lt).__name__ # "LocalTensor" with mode.disable(): regular = torch.ones(2, 2) disabled_type = type(regular).__name__ # "Tensor" return (inside_type, disabled_type), ("LocalTensor", "Tensor")LocalTensorMode.disable()上下文管理器(torch/distributed/_local_tensor/__init__.py)在需要执行「普通张量」逻辑(如构造参考值、打印调试信息)时很有用。对于可移植代码,推荐maybe_disable_local_tensor_mode():无论当前是否处于 LocalTensorMode 中,块内工厂函数都产出普通Tensor(torch/distributed/_local_tensor/__init__.py):
def use_maybe_disable(): """Use maybe_disable_local_tensor_mode() for portable code.""" def create_tensor(): with maybe_disable_local_tensor_mode(): return torch.tensor([1.0, 2.0, 3.0]) t1 = create_tensor() # 普通环境 outside_type = type(t1).__name__ with LocalTensorMode(4): t2 = create_tensor() # 仍处于 mode 内,但强制普通张量 inside_type = type(t2).__name__ return (outside_type, inside_type), ("Tensor", "Tensor")另外,@maybe_run_for_local_tensor装饰器(torch/distributed/_local_tensor/__init__.py)会在 LocalTensorMode 内逐 rank 各执行一次被装饰函数,自动拆解 LocalTensor 输入并聚合逐 rank 输出,适合封装「每个 rank 读自己那一段数据」这类非 SPMD 逻辑,例如按 rank 计算数据偏移切片(见 example_05_rank_specific.py 中的use_maybe_run_decorator)。
7. 示例 6:多维 Mesh 与混合并行
完整源码见 example_06_multidim_mesh.py。2D/3D device mesh 可模拟混合并行(数据并行 DP + 张量并行 TP + 流水线并行 PP)。
创建 2D mesh:
def create_2d_mesh(): """Create a 2D mesh for hybrid parallelism.""" world_size = 8 dp_size, tp_size = 4, 2 with LocalTensorMode(world_size): mesh = init_device_mesh("cpu", (dp_size, tp_size), mesh_dim_names=("dp", "tp")) shape = mesh.shape # (4, 2) dim_names = mesh.mesh_dim_names # ("dp", "tp") total_size = mesh.size() # 8 expected = ((dp_size, tp_size), ("dp", "tp"), world_size) return (shape, dim_names, total_size), expected混合并行(DP + TP)矩阵乘法:
def hybrid_parallelism(): """Combine data parallel and tensor parallel.""" world_size = 8 dp_size, tp_size = 4, 2 with LocalTensorMode(world_size): mesh = init_device_mesh("cpu", (dp_size, tp_size), mesh_dim_names=("dp", "tp")) x = torch.randn(16, 8) dx = distribute_tensor(x, mesh, [Shard(0), Replicate()]) # dp 维切 batch w = torch.randn(8, 12) dw = distribute_tensor(w, mesh, [Replicate(), Shard(1)]) # tp 维切输出通道 dy = dx @ dw expected = x @ w actual = dy.full_tensor().reconcile() return actual, expected3D mesh(DP + TP + PP):
def create_3d_mesh(): """Create a 3D mesh for DP + TP + PP.""" world_size = 24 pp_size, dp_size, tp_size = 2, 3, 4 with LocalTensorMode(world_size): mesh = init_device_mesh( "cpu", (pp_size, dp_size, tp_size), mesh_dim_names=("pp", "dp", "tp"), ) tensor = torch.randn(8, 16, 32) dt = distribute_tensor(tensor, mesh, [Replicate(), Shard(0), Shard(2)]) actual = dt.full_tensor().reconcile() return actual, tensor多维 mesh 下Placements列表长度与 mesh 维度数一致:(pp, dp, tp)各维分别指定Replicate/Shard(dim),聚合后应与原始张量完全一致(测试中使用torch.equal)。
8. 教程示例的自动测试机制
本教程的一个关键工程实践是:文档里的每一段代码都来自源码文件中的函数,测试套件直接调用这些同一批函数,防止文档示例腐化。测试文件为 test_local_tensor_tutorial_examples.py。
各示例函数统一返回(actual, expected)元组,测试只做比对、不含硬编码期望值:
# 摘自 test_local_tensor_tutorial_examples.py from example_01_basic_operations import create_local_tensor def test_create_local_tensor(self): lt, (exp_shape, exp_ranks, exp_rank_0, exp_rank_1) = create_local_tensor() self.assertIsInstance(lt, LocalTensor) self.assertEqual(lt.shape, exp_shape) self.assertEqual(lt._ranks, exp_ranks) self.assertTrue(torch.equal(lt._local_tensors[0], exp_rank_0)) self.assertTrue(torch.equal(lt._local_tensors[1], exp_rank_1))涉及集合通信/DTensor 的测试类在setUpClass中初始化 fake 进程组(不同示例 world_size 分别为 3/4/3/24),tearDownClass中销毁,例如TestExample02CollectiveOperations.setUpClass使用dist.init_process_group("fake", rank=0, world_size=3)。每个示例模块的if __name__ == "__main__":入口也按同样方式初始化 fake 进程组,方便单独运行验证。
9. API 参考
核心实现均位于 torch/distributed/_local_tensor/__init__.py(包路径torch.distributed._local_tensor):
核心类
| 类 | 位置 | 说明 |
|---|---|---|
LocalTensor | L922 | torch.Tensor子类;_local_tensors: dict[rank, Tensor]保存各 rank 分片;主要方法reconcile()(L1124,分片全同时收敛为单张量)、is_contiguous()/contiguous()(逐 rank 处理)、tolist()、numpy() |
LocalTensorMode | L1227 | TorchDispatchMode;构造参数为 world_size 或 rank 集合;disable()(L1477)临时还原普通张量语义;rank_map()(L1499)按 rank 构造分片;tensor_map()(L1508)按 rank 变换分片 |
LocalIntNode | L423 | 逐 rank 整型值容器;_local_ints字典;add/sub/mul/floordiv/mod、sym_max/sym_min/sym_sum、比较运算等 |
工具函数
| 函数 | 位置 | 说明 |
|---|---|---|
local_tensor_mode() | L1794 | 获取当前活跃的LocalTensorMode实例(无则None) |
enabled_local_tensor_mode() | L1811 | 获取已启用的 mode,供库代码感知环境 |
maybe_run_for_local_tensor(func) | L1827 | 装饰器:mode 内逐 rank 执行被装饰函数并聚合输出 |
rank_map(cb) | L1862 | 函数版rank_map:rank -> Tensor回调生成逐 rank 值 |
tensor_map(tensor, cb) | L1877 | 函数版tensor_map:(rank, shard) -> Tensor回调逐 rank 变换 |
maybe_disable_local_tensor_mode() | L1898 | 上下文管理器:块内保证产出普通张量,是否真正禁用取决于当前环境 |
10. 最佳实践与常见陷阱
最佳实践:
- 仅用于测试:
LocalTensor开销显著(逐 rank 分发执行),不要用于生产代码。 - 初始化进程组:即便只做本地测试,涉及集合通信也需初始化进程组(使用
"fake"后端,如dist.init_process_group("fake", rank=0, world_size=N))。 - 避免在内部张量上设置 requires_grad:
LocalTensor要求内部各分片张量requires_grad=False,需要在LocalTensor包装层上设置梯度。 - 断言用 reconcile():当所有 rank 应当具有相同值时(例如 all-reduce 之后),用
reconcile()收敛出单个张量再做断言。 - 调试时直接访问分片:通过
tensor._local_tensors[rank](或tensor._local_tensor_<rank>属性)检查单个 rank 的状态。
常见陷阱:
- 忘记上下文管理器:在
LocalTensorMode之外对 LocalTensor 的算子仍然可以工作,但工厂函数(torch.ones等)不会再自动创建 LocalTensor。 - rank 不匹配:同一操作中参与运算的各 LocalTensor 的 rank 集合必须兼容,构造分片字典时注意 rank 键一致。
- 内部张量带梯度:用
requires_grad=True的张量构造 LocalTensor 会抛错;请在包装层处理梯度。
小结
LocalTensor为 PyTorch 分布式开发提供了一种「进程内 SPMD 模拟器」:以{rank: shard}字典语义承载各 rank 状态,借助LocalTensorMode让常规张量 API 无感产出 LocalTensor,再配合 DTensor 的distribute_tensor/init_device_mesh覆盖数据并行、张量并行、混合并行乃至非均匀分片等场景。所有示例均有对应的可执行源码(test/distributed/local_tensor_tutorial_examples/)与测试(test_local_tensor_tutorial_examples.py)保障正确性,是本地开发与 CI 中调试分布式张量代码的实用工具。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考