mindspore-ai / mindspore-ai/hyper-parallel

[笔记] PyTorch TensorImpl、Storage 与 HyperParallel DTensor 创建机制

Open
#188 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

PyTorch TensorImpl、Storage 与 HyperParallel DTensor 创建机制笔记

1. 这份笔记回答什么

本文记录以下几个容易混淆的问题:

  • Python 层的 Tensor 对象、C++ 层的 TensorImpl 和实际数据 Storage 分别是什么;
  • 为什么不同 Tensor 对象可以共享同一份底层内存;
  • tensor.data 为什么与 tensor 的 Python id、TensorImpl 不同,但通常共享 Storage;
  • HyperParallel 的 Torch DTensor 如何通过 Tensor._make_subclass() 创建;
  • DTensor 本体、DTensor._local_tensorDTensor.data 之间是什么关系;
  • nn.Parameter(DTensor) 为什么又会产生一层新的 TensorImpl;
  • meta init / to_empty() 为什么可能让 DTensor 本体与 _local_tensor 的 Storage 暂时失配;
  • 修复这种失配时,为什么不能直接交换 _local_tensor 本体。

本文是机制备忘录,不定义新的用户 API,也不讨论完整的 DTensor 分布式语义。

2. PyTorch Tensor 的三层模型

可以把一个 PyTorch Tensor 简化成三层:

Python Tensor 对象
    │
    │ 持有 C++ Tensor handle
    ▼
TensorImpl
    │
    │ 持有 Storage 引用和 view metadata
    ▼
Storage / StorageImpl
    │
    ▼
CPU / NPU / GPU 实际内存
2.1 Python Tensor 对象

这是 Python 代码中看到的对象:

tensor = torch.ones(4)

id(tensor) 标识的是 Python 对象身份。模块参数注册、优化器保存的参数引用以及用户持有的变量,首先依赖这一层的对象身份。

两个不同的 Python Tensor 对象可以持有不同的 TensorImpl,也可以通过这些 TensorImpl 共享同一个 Storage。

2.2 TensorImpl

TensorImpl 是 Tensor 的核心 C++ 元数据对象。简化来看,它记录:

  • 对 Storage / StorageImpl 的引用;
  • storage_offset
  • sizes 和 strides;
  • dtype;
  • device 和 dispatch keys;
  • autograd metadata、version counter 等运行时状态。

因此,整体替换 TensorImpl 会同时改变当前 Tensor 看到的 Storage、设备、dtype、shape、stride 等信息。

TensorImpl 自身不是实际数据内存。多个 TensorImpl 可以引用同一个 Storage,并用不同的 offset、sizes 和 strides 表达不同 view。

2.3 Storage / StorageImpl

Storage 管理实际内存区域。简化来看,它包含:

  • data pointer;
  • 分配字节数;
  • device;
  • allocator;
  • 是否允许 resize 等属性。

常用观测方式:

storage = tensor.untyped_storage()
storage_id = storage._cdata       # StorageImpl 标识,仅用于调试
memory_ptr = storage.data_ptr()   # 实际内存地址;meta/空 tensor 可能为 0

判断“是否为同一份底层存储”时,StorageImpl 标识比只比较 data_ptr() 更严格。两个不同 view 可以共享同一个 StorageImpl,但 tensor.data_ptr() 还会受到 storage_offset 影响。

3. 多个 TensorImpl 共享一个 Storage

典型 view 关系如下:

Tensor A ──> TensorImpl A ──┐
                            ├──> Storage X
Tensor B ──> TensorImpl B ──┘

两个 Tensor:

  • Python 对象不同;
  • TensorImpl 不同;
  • StorageImpl 相同;
  • 可能拥有不同的 shape、stride 和 storage offset。

对应 Mermaid:

flowchart LR
    A[Python Tensor A] --> IA[TensorImpl A]
    B[Python Tensor B] --> IB[TensorImpl B]
    IA --> S[StorageImpl X]
    IB --> S
    S --> M[Device Memory]

4. tensor.data 的读取机制

读取:

tensor_data = tensor.data

不是读取一个长期保存在 tensor.__dict__ 中的普通属性。Tensor.data 是 descriptor;getter 会返回一个 detached alias。

通常关系为:

tensor      ──> TensorImpl A ──┐
                               ├──> Storage X
tensor.data ──> TensorImpl D ──┘

所以一般有:

id(tensor) != id(tensor.data)
tensor._cdata != tensor.data._cdata
tensor.untyped_storage()._cdata == tensor.data.untyped_storage()._cdata

每次读取 .data 都可能产生新的临时 Tensor 对象和 TensorImpl:

data1 = tensor.data
data2 = tensor.data

assert id(data1) != id(data2)
assert data1._cdata != data2._cdata
assert data1.untyped_storage()._cdata == data2.untyped_storage()._cdata

.data 会绕开正常 autograd 语义,生产逻辑中应谨慎使用;这里主要用它解释 alias 和底层存储关系。

5. tensor.data = another_tensor 与普通 Python 赋值不同

赋值:

tensor.data = another_tensor

会调用 Tensor.data descriptor 的 setter,而不是让一个名为 data 的普通 Python 属性指向 another_tensor

简化后的效果是:保留 tensor 的 Python 对象身份,并修改其底层数据相关状态,使其引用 another_tensor 的数据:

赋值前:tensor ──> TensorImpl A ──> Storage X
赋值后:tensor ──> TensorImpl A ──> Storage Y

实际 setter 还会做 Tensor 类型、设备和 autograd 兼容性检查。对于 Tensor subclass,普通 Tensor 与 subclass 之间可能被判定为 incompatible tensor type;meta 与实体设备之间也不能简单依赖 set_() 跨设备重绑 Storage。

这与交换两个 .data 临时 alias 不同:

torch.utils.swap_tensors(tensor.data, another_tensor.data)

这里只改变 getter 返回的两个临时 Python Tensor 对象,不会回写 tensor 本体所持有的 TensorImpl。

6. HyperParallel Torch DTensor 的创建链路

关键代码位置:

  • hyper_parallel/core/dtensor/dtensor.py::DTensor.from_local()
  • hyper_parallel/platform/torch/dtensor.py::DTensorBase.__new__()
  • hyper_parallel/core/dtensor/dtensor.py::DTensor.__init_data__()

调用链:

DTensor.from_local(local_tensor, mesh, placements)
    │
    ├── 构建或复用 Layout
    │
    └── DTensor(local_tensor, mesh, placements, layout)
            │
            ├── DTensorBase.__new__
            │     └── Tensor._make_subclass(cls, local_tensor, requires_grad)
            │
            └── __init_data__
                  ├── self._local_tensor = local_tensor
                  ├── self._device_mesh = device_mesh
                  ├── self._layout = layout
                  └── self._placements = placements

Mermaid 表示:

flowchart TD
    L[local_tensor] --> F[DTensor.from_local]
    F --> B[build or reuse Layout]
    F --> N[DTensorBase.__new__]
    N --> M[Tensor._make_subclass]
    M --> D[DTensor Python object and TensorImpl]
    D --> I[DTensor.__init_data__]
    I --> R[store local_tensor as _local_tensor]
    I --> E[store mesh layout placements]
6.1 _make_subclass() 与 local tensor 的关系

核心调用:

t = Tensor._make_subclass(
    cls,
    local_tensor,
    local_tensor.requires_grad,
)

它创建一个 Tensor subclass 对象及其 TensorImpl。该 TensorImpl 与传入的 local_tensor TensorImpl 不同,但两者共享 local tensor 的 StorageImpl。

随后:

t.__init_data__(local_tensor, device_mesh, placements, layout)

又将传入的 Python local tensor 对象保存到 t._local_tensor

裸 DTensor 创建后的长期关系为:

DTensor 本体          ──> TensorImpl D ──┐
                                         ├──> Local Storage X
DTensor._local_tensor ──> TensorImpl L ──┘

其中 DTensor 本体和 _local_tensor 是不同的 Python 对象、不同的 TensorImpl,但共享同一个 StorageImpl。

访问 DTensor.data 后,还会多出一个临时 alias:

DTensor 本体          ──> TensorImpl D ──┐
DTensor._local_tensor ──> TensorImpl L ──┼──> Local Storage X
DTensor.data          ──> TensorImpl A ──┘

7. HyperParallel 的 DTensor.data override

Torch 后端在 DTensorBase 中覆盖了 .data

getter 的目标是绕过 __torch_function__,直接读取 DTensor 本体对应的 Tensor data alias:

with torch._C.DisableTorchFunctionSubclass():
    return Tensor.data.__get__(self, type(self))

这里的 DisableTorchFunctionSubclass 很重要。如果不绕过 subclass dispatch,观测结果可能被 HyperParallel DTensor 分发逻辑改写为 _local_tensor 路径,无法准确看到 DTensor 本体所持有的 TensorImpl/Storage。

setter 的意图是同时更新两条长期引用:

Tensor.data.__set__(self, local_value)
Tensor.data.__set__(self._local_tensor, local_value)

即保持:

DTensor body Storage == DTensor._local_tensor Storage

但底层 PyTorch 对 Tensor subclass、跨设备及 meta materialization 有兼容性限制,不能假定任意场景都能通过 .data = ... 完成修复。

8. nn.Parameter(DTensor) 的创建机制

PyTorch 对 custom Tensor subclass 采用 Parameter 的特殊路径:

t = data.detach().requires_grad_(requires_grad)
t._is_param = True
return t

HyperParallel DTensor.detach() 会:

detached_local = self._local_tensor.detach()
return self.__class__(
    detached_local,
    device_mesh=self._device_mesh,
    placements=self._alias_placements(),
)

因此 nn.Parameter(DTensor) 并不是简单给原 DTensor Python 对象加一个标签。它会经过:

原 DTensor
    │
    └── detach()
          ├── 创建 detached local tensor
          └── 再次调用 DTensor 构造
                └── 再次调用 Tensor._make_subclass()

最终 Parameter DTensor 仍然满足:

sharded_param 本体          ──> TensorImpl P ──┐
                                               ├──> Sharded Storage X
sharded_param._local_tensor ──> TensorImpl L ──┘

访问 sharded_param.data 时,再创建临时 TensorImpl A,同样共享 Storage X。

flowchart LR
    D[DTensor] --> X[DTensor.detach]
    X --> LD[detached local tensor]
    LD --> C[DTensor constructor]
    C --> MS[Tensor._make_subclass]
    MS --> P[DTensor marked as Parameter]
    P --> SP[sharded_param TensorImpl]
    P --> LP[_local_tensor TensorImpl]
    SP --> S[shared sharded Storage]
    LP --> S

9. Deferred init 为什么会打破 Storage 一致性

FSDP meta-init 场景中,初始 sharded DTensor 可能位于 meta device:

sharded_param 本体          ──> TensorImpl P ──┐
                                               ├──> Meta Storage
sharded_param._local_tensor ──> TensorImpl L ──┘

module.to_empty(device="npu") 物化模块参数后,模块上的参数可能已成为实体设备上的普通 Tensor/Parameter,而 HSDP state 仍保存原来的 sharded DTensor 对象身份。

reset_sharded_param() 会根据物化参数构建新的 local view,并执行:

self.sharded_param._local_tensor = local_view

此时只替换了 Python 属性 _local_tensor 指向的对象,没有更新 sharded DTensor 本体持有的 TensorImpl:

修复前暂态:

sharded_param 本体          ──> TensorImpl P ──> Meta Storage
sharded_param.data          ──> TensorImpl A ──> Meta Storage
sharded_param._local_tensor ──> TensorImpl L ──> NPU Storage X

这里是三份 TensorImpl、两份 StorageImpl。DTensor 的部分 Python property 会从 _local_tensor 返回 device/dtype,看起来已经在 NPU;但直接使用 DTensor 本体 TensorImpl 的底层路径仍可能看到 meta Storage。这种“表层元数据已物化、本体 TensorImpl 未物化”的状态必须修复。

10. TensorImpl 修复的实际机制

当前讨论的精简修复思路为:

sharded_param_data = self.sharded_param.data
local_tensor_data = self.sharded_param._local_tensor.data

if storage_mismatch:
    local_tensor_data.requires_grad_(self.sharded_param.requires_grad)
    torch._C._swap_tensor_impl(self.sharded_param, local_tensor_data)

关键点:参与交换的是 _local_tensor.data 临时 alias,不是 _local_tensor 本体。

10.1 swap 前
sharded_param            ──> TensorImpl P(meta) ──> Meta Storage
_local_tensor            ──> TensorImpl L      ──> NPU Storage X
local_tensor_data alias  ──> TensorImpl A      ──> NPU Storage X
10.2 只交换 TensorImpl 后
sharded_param            ──> TensorImpl A      ──> NPU Storage X
_local_tensor            ──> TensorImpl L      ──> NPU Storage X
local_tensor_data alias  ──> TensorImpl P(meta)──> Meta Storage

临时 alias 在函数结束后释放,旧 meta TensorImpl 随之失去该临时引用。_local_tensor 从未参与 swap,所以不会变成 meta。

flowchart TB
    subgraph Before[Before repair]
        BP[sharded_param] --> BI[Meta TensorImpl]
        BL[_local_tensor] --> BLI[Local TensorImpl]
        BA[local_tensor.data alias] --> BAI[Alias TensorImpl]
        BI --> BM[Meta Storage]
        BLI --> BN[NPU Storage X]
        BAI --> BN
    end

    subgraph After[After TensorImpl swap]
        AP[sharded_param] --> AAI[Alias TensorImpl]
        AL[_local_tensor] --> ALI[Local TensorImpl]
        AA[temporary alias] --> AMI[Old Meta TensorImpl]
        AAI --> AN[NPU Storage X]
        ALI --> AN
        AMI --> AM[Meta Storage]
    end

11. 为什么不能交换两个 .data alias

下面的代码无效:

torch.utils.swap_tensors(
    self.sharded_param.data,
    self.sharded_param._local_tensor.data,
)

因为两个参数都是 getter 返回的临时 Python Tensor 对象。交换只影响这两个临时对象,不会改变 self.sharded_param 本体持有的 TensorImpl。

12. 为什么不能直接 swap _local_tensor 本体

下面的逻辑会污染 _local_tensor

swap_impl(self.sharded_param, self.sharded_param._local_tensor)

交换后 _local_tensor 会拿到旧 meta TensorImpl:

sharded_param -> NPU TensorImpl
_local_tensor -> Meta TensorImpl  # 错误

所以必须使用 _local_tensor.data 生成的临时 alias 作为 donor。alias 与 _local_tensor 共享 NPU Storage,但拥有独立 TensorImpl;旧 meta TensorImpl 最终落在 alias 上,而不是 _local_tensor 上。

13. torch.utils.swap_tensors()_swap_tensor_impl() 的差异

公开接口 torch.utils.swap_tensors(t1, t2) 不只交换 TensorImpl,还会交换:

  • __class__
  • __dict__
  • slots;
  • 底层 TensorImpl。

因此不能直接用普通 Tensor alias 与 DTensor Parameter 本体调用:

torch.utils.swap_tensors(sharded_param, local_tensor_data)

否则 sharded_param 的 DTensor/Parameter Python 壳和 mesh/layout 等动态属性也会被交换给临时 alias。

如果只希望保留原 Python 对象、Parameter 身份和 DTensor 元数据,仅替换其 TensorImpl,就需要更窄的 TensorImpl 级操作。但 _swap_tensor_impl() 是 PyTorch 私有接口,存在版本兼容风险;使用时应有定向测试覆盖,或者使用同类型、同 slots、完整元数据的 DTensor donor 配合公开 swap_tensors()

14. 调试与验证方法

14.1 建议观测四类标识
def inspect_tensor(tensor):
    with torch._C.DisableTorchFunctionSubclass():
        storage = torch.Tensor.untyped_storage(tensor)
        tensor_impl = torch.Tensor._cdata.__get__(tensor, torch.Tensor)
    return {
        "python_id": id(tensor),
        "tensor_impl": tensor_impl,
        "storage_impl": storage._cdata,
        "data_ptr": storage.data_ptr(),
        "storage_device": storage.device,
        "offset": tensor.storage_offset(),
        "size": tuple(tensor.size()),
        "stride": tuple(tensor.stride()),
    }

对 DTensor 调试时需要绕过 __torch_function__。HyperParallel 的 devicedtype 等 property 可能主动返回 _local_tensor 的信息,仅打印 dtensor.device 不足以证明 DTensor 本体 TensorImpl 已经物化。

14.2 验证同 Storage 时不要只看 data_ptr()

建议同时比较:

lhs.untyped_storage()._cdata == rhs.untyped_storage()._cdata
lhs.untyped_storage().data_ptr() == rhs.untyped_storage().data_ptr()
lhs.storage_offset() == rhs.storage_offset()
lhs.size() == rhs.size()
lhs.stride() == rhs.stride()

原因:

  • meta tensor 和部分空 tensor 的 data_ptr() 都可能为 0;
  • 同一 Storage 的不同 view,tensor.data_ptr() 可能因 offset 不同而不同;
  • 指针相同并不自动意味着 view metadata 相同。

15. 本次实测结论

在 HyperParallel Torch DTensor 上观测到:

  1. 裸 DTensor 创建后,DTensor 本体与 _local_tensor 是不同 Python 对象、不同 TensorImpl,共享同一个 StorageImpl;
  2. 每次读取 DTensor.data 会得到新的临时 Python Tensor/TensorImpl,并继续共享 DTensor 本体当时的 Storage;
  3. nn.Parameter(DTensor) 会通过 DTensor 的 detach() 再构造一个 DTensor,并设置 _is_param
  4. 正常 sharded Parameter 状态下,Parameter DTensor 本体、其 _local_tensor 以及当次读取的 .data alias 共享同一个 StorageImpl;
  5. deferred init 只更新 _local_tensor 时,会暂时形成 DTensor 本体为 meta Storage、_local_tensor 为 NPU Storage 的失配;
  6. 使用 _local_tensor.data 临时 alias 作为 TensorImpl donor,可以在不污染 _local_tensor 的前提下让 DTensor 本体重新引用 NPU Storage;
  7. 直接交换两个 .data alias 不会修改父 Tensor;直接交换 _local_tensor 本体则会把旧 meta TensorImpl 交换给 _local_tensor

16. 相关代码位置

  • hyper_parallel/platform/torch/dtensor.py
    • DTensorBase.__new__()
    • DTensorBase.detach()
    • DTensorBase.data
  • hyper_parallel/core/dtensor/dtensor.py
    • DTensor.from_local()
    • DTensor.__init_data__()
  • hyper_parallel/platform/torch/fully_shard/param.py
    • TorchHSDPParamV2._resolve_reset_param()
    • TorchHSDPParamV2.reset_sharded_param()
    • TorchHSDPParamV2._update_shardedparam_storage_forcely()

17. 一句话记忆

Python Tensor 是对象身份,TensorImpl 是 Tensor 的数据与运行时元数据入口,StorageImpl 才管理实际内存;HyperParallel DTensor 用独立 TensorImpl 包装 local tensor 并共享其 Storage,同时把 local tensor 作为 _local_tensor 长期保存。

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 331
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/331

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reading DTensor.from_local() and DTensor.init_data() in hyper_parallel/core/dtensor/dtensor.py, then DTensorBase.new() in hyper_parallel/platform/torch/dtensor.py. Verify the note’s TensorImpl, Storage, _make_subclass(), Parameter, and meta-init explanations against these entry points. Done means the mechanism note is added accurately and its documented relationships remain consistent.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
documentation
Issue type
Documentation
Difficulty
4/5
Estimated time
3-5 days
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
68/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.