mindspore-ai / mindspore-ai/hyper-parallel

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

Open
#182 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

模型高性能实现替换: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 保持现状。

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

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 in components/model_transform/replacement.py, especially the compile and apply entry points, then inspect trainer/config.py and the Transformers conversion-mapping APIs. Add replacement unit coverage for type metadata, is_fusable, scoped transforms, aliases, conflicts, atomic failure, and mapping visibility. Done means fused parameters load through get_model_conversion_mapping(model) while the checkpoint loader remains unchanged.

Written by the indexing model from the issue text.

Assessment

Tech stack
huggingface, python, pytorch
Domain
distributed-systems, machine-learning
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.