已开启
模型高性能实现替换:`weights_mapping` 接入 #338
TonghanZhang创建于  8月17日
TonghanZhang
TonghanZhang成员
8月17日 创建

模型高性能实现替换:weights_mapping 接入

1. 当前代码基线

新代码已经实现 PR #1156 的后构造 module replacement,真实调用链是:

parallelism.plan_overrides
→ PlanOverride
→ entries_to_module_replacements()
→ compile_module_replacements()
→ apply_module_replacements()
→ _apply_module_replacement_actions()
→ apply_model_infrastructure()
→ sharding / FSDP2
→ CheckpointManager.load_checkpoint()

当前 apply_module_replacements() 接口直接返回 model:

apply_module_replacements(
    model: nn.Module,
    plan: ModuleReplacementPlan,
    *,
    context: Mapping[str, Any] | None = None,
) -> nn.Module

replacement target 使用 @module_replacement 声明,接收 modulemodule_fqn、只读 context,当前必须构造并返回 nn.Moduleapply_module_replacements() 会原地更新父模块的 _modules,最后返回同一个根 model。

当前 _validate_replacement() 只允许结构保持替换:

  • module、parameter、buffer 注册名不变;
  • parameter 和 buffer 对象 identity 不变;
  • state_dict key 不变;
  • forward 调用协议、train/eval 状态不变。

PR #1156 已覆盖 RMSNorm 等“换实现、不换参数 schema”的场景。gate_proj + up_proj → gate_up_proj 这类参数融合需要在现有 replacement executor 上增加 weights_mapping

2. 用户接口

YAML 只选择实例位置和 replacement target:

parallelism:
  plan_overrides:
    - match: "model.layers.*.input_layernorm"
      replace_module:
        _target_: perf_kernels.NpuRMSNorm

    - match: "model.layers.*.mlp"
      replace_module:
        _target_: perf_modules.FusedMLP

原 module 的兼容类型和参数转换关系由 replacement 实现维护。match 负责选择实例,replace_module._target_ 负责选择实现;YAML 无需引用 Transformers 内部 class path。

普通函数和自定义 autograd 函数由 replacement module 的 forward 调用。直接函数替换属于 codegen 稳定调用点对应的独立能力。

3. replacement 声明

3.1 类型约束由 target 声明

Hyper 扩展 @module_replacement,由 replacement class 声明可接受的原 module 类型,并实现完整的 module 生命周期:

@module_replacement(module_type=LlamaRMSNorm)
class NpuRMSNorm(nn.Module):
    def __init__(self, *, module, module_fqn, context):
        super().__init__()

        # 复用同一个 Parameter,保持参数 identity 和 state_dict key。
        self.weight = module.weight
        self.variance_epsilon = module.variance_epsilon
        self.train(module.training)

    def forward(self, hidden_states):
        return npu_rms_norm(
            hidden_states,
            self.weight,
            self.variance_epsilon,
        )

当前 match 遍历 model.named_modules(),命中的是 model tree 中已经注册的 nn.Module FQN。因此 target 也必须构造 nn.Module,以承接原模块的 Parameter、buffer、training state、hooks、state_dictforward 协议。

npu_rms_norm 是计算函数,不是 model tree 节点:它没有 module FQN,也不持有 Parameter,当前 executor 无法用 module FQN 直接替换它。正确路径是:

match input_layernorm module
→ replacement target 构造 NpuRMSNorm
→ NpuRMSNorm.forward() 调用 npu_rms_norm()

直接 target 到 npu_rms_norm 需要 codegen 暴露函数调用点,并由另一套 call-site replacement 机制处理;它不属于当前 module replacement 协议。

module_type 支持一个 nn.Module class 或 class tuple;exact_type=False 默认使用 isinstance(),需要排除子类时由 decorator 设置 exact_type=True。这些元数据写入 target class,entries_to_module_replacements() 读取后构造 ModuleReplacementSpec。类型兼容性只维护在 target class 中。需要超出类型检查的结构约束时,target class 提供 @classmethod is_fusable(cls, module: nn.Module) -> bool。输入与 Transformers ModuleFusionSpec.is_fusable(module) 相同;未声明该方法时,类型检查通过即视为兼容。

3.2 make_transforms(config) 由 replacement class 声明

replacement class 的构造函数生成 nn.Module 实例;同一个实例通过与 Transformers 一致的接口描述参数转换:

@module_replacement(module_type=LlamaMLP)
class FusedMLP(nn.Module):
    def __init__(self, *, module, module_fqn, context):
        super().__init__()
        ...

    def make_transforms(
        self,
        config: "PretrainedConfig",
    ) -> list[WeightTransform]:
        # gate/up 融合本身不依赖 config,但统一协议仍接收根模型 config。
        return [
            WeightConverter(
                source_patterns=[
                    "gate_proj.weight",
                    "up_proj.weight",
                ],
                target_patterns="gate_up_proj.weight",
                operations=[Concatenate(dim=0)],
            ),
        ]

config 的来源固定为根 PreTrainedModel.config,与 Transformers ModuleFusionSpec.make_transforms(config) 的输入相同。需要模型配置的转换直接读取该对象;例如 Transformers 的 patch-embedding fusion 从 vision_config 读取 patch_sizetemporal_patch_sizein_channels

mapping 描述 replacement class 的参数 schema,因此与 class 定义放在一起;pattern 使用相对于命中 module 的参数名。参数 schema 保持不变的 class 省略该方法;executor 将方法不存在规范化为空 transforms。每次调用都创建新的 transform,因为 WeightTransform 会记录匹配和加载状态。

4. 编译与执行

4.1 compile_module_replacements()

compile 保持 PR #1156 的职责,只做匹配和静态校验:

  1. match 查找 module 和全部 alias FQN;
  2. 使用 target class 声明的 module_typeexact_type 校验原 module;
  3. target class 提供 is_fusable(module) 时执行额外结构兼容性检查;
  4. 返回不可变 ModuleReplacementPlan

4.2 apply_module_replacements()

apply 在现有 replacement 构造流程中同时收集 transforms:

apply_module_replacements(
    model: nn.Module,
    plan: ModuleReplacementPlan,
    *,
    context: Mapping[str, Any] | None = None,
) -> nn.Module

执行顺序:

  1. 构造全部 replacement module;
  2. 校验 replacement 的返回类型、train/eval 状态、forward 和 hook;
  3. 调用可选的 replacement.make_transforms(model.config)
  4. 为每个相对 transform 创建独立副本并设置实例作用域;
  5. 校验新旧参数 schema 以及 source/target 冲突;
  6. 确认全部 source module 仍注册在原 FQN;
  7. 安装 replacement;
  8. 将新增 transforms 合并到 Transformers 原生 conversion mapping。

实例作用域使用 WeightTransform 已有字段:

transform.scope_prefix = matched_module_fqn
transform.base_model_prefix = model.base_model_prefix

共享 module 的多个注册 FQN 为每个可加载路径生成独立 transform。

4.3 更新 Transformers 原生 mapping

get_model_conversion_mapping(model) 是读取接口;修改它返回的临时 list 不会更新 registry。apply 使用 Transformers 的注册接口更新 root model mapping:

existing = extract_weight_conversions_for_model(model) or []
merged = [*existing, *extra_transforms]

register_checkpoint_conversion_mapping(
    type(model).__name__,
    merged,
    overwrite=True,
)

这里使用 extract_weight_conversions_for_model(model) 只取得 root model 的 class 或 model_type mapping。直接把 get_model_conversion_mapping(model) 的完整结果重新注册,会把 nested submodel 和 legacy transforms 一起注册到 root,后续读取时产生重复。

apply_module_replacements() 完成注册后仍只返回 model,现有外层调用保持不变:

plan = compile_module_replacements(model, entries_to_module_replacements(entries))
return apply_module_replacements(model, plan)

后续 get_model_conversion_mapping(model) 会通过 root class name 取得合并后的 transforms,checkpoint loader 无需修改。

该 registry 是进程级全局状态。同一进程内,同一个 root model class 的所有实例共享 mapping;后一次注册会影响后续构造或加载的同 class model。若必须支持同 class 实例使用不同 replacement YAML,Transformers 原生 registry 无法提供实例隔离,届时才需要实例级 mapping 传递路径。

5. 校验边界

make_transforms(config) 校验
未实现或返回空列表 保持 PR #1156 的严格校验:注册名、identity、state_dict key 全部不变
返回非空列表 允许 parameter/buffer schema 改变;仍校验返回值、train/eval 状态、forward 兼容性、hook 限制和 transform 完整性

有 transforms 时至少检查:

  1. pattern 仅描述当前 module 的相对 key;
  2. 原 module 消失的持久化 key 必须由 source pattern 覆盖;
  3. replacement 新增的持久化 key 必须由 target pattern 覆盖;
  4. 未参与转换的同名 parameter/buffer 仍保持原对象 identity;
  5. 多条规则展开后,converter source 冲突或同一 target 被重复生成时立即报错;
  6. 全部 replacement、transforms 和校验完成后才安装 replacement。

make_transforms() 的返回值必须是 list[WeightTransform],具体 transform 的构造和反向转换语义沿用 Transformers 定义。

6. 实现改动位置

文件 改动
components/model_transform/replacement.py 扩展 decorator 的 source type 元数据和可选 is_fusable;apply 阶段构造 replacement、调用可选 make_transforms(model.config)、校验 scoped transforms,并注册到 Transformers root model mapping
trainer/config.py YAML replacement 不再解析 module_type/exact_type;从 target decorator 元数据构造 spec
replacement 单测 覆盖类型声明、is_fusable、无 transforms、融合 transforms、多个 FQN、alias、冲突、原子失败及 get_model_conversion_mapping(model) 可见性

7. Model transform 验收

原 module:
  gate_proj.weight
  up_proj.weight

YAML:
  match = model.layers.*.mlp
  replace_module = FusedMLP

compile 输出:
  plan.targets = 命中的原 module

apply 输出:
  model.layers.0.mlp = FusedMLP
  model.layers.0.mlp.gate_up_proj.weight

get_model_conversion_mapping(model):
  WeightConverter(
    source = gate_proj.weight + up_proj.weight
    target = gate_up_proj.weight
    scope_prefix = model.layers.0.mlp
  )

验收覆盖 replacement class 构造、make_transforms(model.config) 调用、FQN scope、schema 覆盖、冲突检查、原子安装,以及新增 transforms 能由 Transformers 原生 get_model_conversion_mapping(model) 读取;checkpoint loader 保持现状。

likedislike
TonghanZhangTonghanZhang成员
8月17日 修改了issue 的描述