Skip to content

Module::To 在设备或数据类型转换时破坏跨子模块的共享参数关系 #236

Description

@superme1one

问题描述

当不同子模块引用同一个 Tensor 参数时,调用 Module::To(Device) 或 Module::To(DataType) 可能破坏原有的参数共享关系。

本次使用一个小型 GPT-2 模型复现。模型的词嵌入层与语言模型输出层原本共享同一个权重张量:

  • CPU 迁移到 GPU 后,两个参数变成不同的 Tensor 对象,底层存储也不再共享。
  • 即使只在 CPU 上进行 FP32 到 FP32 的同类型转换,也会产生不同的 Tensor 对象;不过在该场景下,底层存储仍然共享。
  • 转换后,参数去重结果发生变化,同一个共享权重被统计了两次。

测试环境

  • InfiniTrain 提交:a8b2d3776f04a9eee5c0dbff74b76fdf6c4d5c9a
  • CPU 复现:FP32 到 FP32 的同类型转换。
  • GPU 复现:在天数 Iluvatar TG-V200 的本地 CoreX CUDA 兼容适配版本上进行 CPU 到 GPU 的迁移。
  • 尚未在 NVIDIA GPU 上复现验证。

由于 CPU 同类型转换也能复现,因此该问题不依赖天数 GPU 的本地适配。

复现步骤

创建具有以下配置的小型 GPT-2 模型:

vocab_size = 32
original_vocab_size = 32
n_layer = 1
n_head = 2
n_kv_head = 2
n_embd = 16
block_size = 16

转换前,检查词嵌入层与输出层的权重是否共享:

auto embedding =
    model->module("transformer").module("wte").parameter("weight");

auto lm_head =
model->module("lm_head").parameter("weight");

bool same_tensor = embedding.get() == lm_head.get();
bool same_storage = embedding->DataPtr() == lm_head->DataPtr();

统计去重后的参数元素总数:

size_t parameter_count = 0;
for (const auto &parameter : model->Parameters()) {
parameter_count += parameter->NumElements();
}

然后分别在新建模型上执行以下任一操作:

// CPU 复现:模型原本已经是 FP32。
model->To(DataType::kFLOAT32);

或者:

// 设备迁移复现。
model->To(Device(Device::DeviceType::kCUDA, 0));

转换完成后,重新从模型中获取这两个注册参数,再次检查 Tensor 对象、底层存储和去重后的参数元素总数。

预期结果

设备或数据类型转换应保留原有的参数共享关系。

真正的设备或类型变化可能需要创建新 Tensor,但原来引用同一个参数的位置,应继续引用同一个转换后的 Tensor。

去重后的参数元素总数应保持为 4080。

实际结果

状态 是否为同一 Tensor 对象 是否共享底层存储 去重后的参数元素总数
转换前 是 是 4080
原版 CPU → GPU 否 否 4592
原版 CPU FP32 → FP32 否 是 4592
本地修复后,上述两种场景 是 是 4080

增加的 512 个参数元素对应 32 × 16 的词嵌入/输出层共享权重,被重复统计了一次。

需要特别区分:CPU 同类型转换场景并未证明底层数据被复制,而是 Tensor 对象身份发生了变化,底层存储仍然共享。

影响

共享参数属于模型定义的一部分,转换操作破坏共享关系可能改变训练行为。

对于设备迁移场景,词嵌入层和输出层变成独立权重,后续训练可能分别更新,无法继续保持权重绑定。

对于 CPU 同类型转换场景,不同 Tensor 对象共享同一底层存储,也可能影响梯度归属和优化器参数去重。目前已确认的是对象身份变化和参数重复计数;尚未单独证明该 CPU 场景发生了重复优化器更新,因此不将其作为已经确认的数值结果。

原因分析

在 infini_train/src/nn/modules/module.cc 中,转换操作对每个注册参数分别创建新的 Tensor 包装对象:

std::make_shared<Tensor>(param->To(device))

或者:

std::make_shared<Tensor>(param->To(dtype))

随后递归转换各个子模块,但整个转换过程没有共享一个覆盖完整模块图的转换映射。

因此,即使两个子模块原本引用同一个 Tensor,也可能分别生成不同的转换结果。参数去重依赖 Tensor 对象身份,原来的共享权重便会被当作两个参数。

建议修复方法

在一次根模块转换中,共享统一的转换上下文:

  • 建立“原 Tensor 对象身份 → 转换后 Tensor”的映射。
  • 再次遇到同一个原 Tensor 时,复用已有转换结果。
  • 在整个转换过程中保留原 Tensor 的所有权,避免对象提前释放或地址复用。
  • 记录已访问的模块,避免重复注册的子模块被多次转换。
  • 对共享 buffer 采用相同处理方式。
  • 对真正不需要转换的操作保留原 Tensor 对象。
  • 安全处理空参数、空 buffer 和空子模块注册项。
  • 避免仅因父模块的设备元数据已经匹配目标,就跳过仍需要转换的子模块。

这里的修复范围是:保留对同一个注册 Tensor 对象的共享引用。不同 Tensor 对象之间任意视图或底层存储别名关系,需要另外设计处理,不能直接认为已被覆盖。

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions