mindspore-ai / mindspore-ai/hyper-parallel
模型高性能实现替换:`weights_mapping` 接入
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 声明,接收 module、module_fqn、只读 context,当前必须构造并返回 nn.Module。apply_module_replacements() 会原地更新父模块的 _modules,最后返回同一个根 model。
当前 _validate_replacement() 只允许结构保持替换:
- module、parameter、buffer 注册名不变;
- parameter 和 buffer 对象 identity 不变;
state_dictkey 不变;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_dict 和 forward 协议。
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_size、temporal_patch_size 和 in_channels。
mapping 描述 replacement class 的参数 schema,因此与 class 定义放在一起;pattern 使用相对于命中 module 的参数名。参数 schema 保持不变的 class 省略该方法;executor 将方法不存在规范化为空 transforms。每次调用都创建新的 transform,因为 WeightTransform 会记录匹配和加载状态。
4. 编译与执行
4.1 compile_module_replacements()
compile 保持 PR #1156 的职责,只做匹配和静态校验:
- 按
match查找 module 和全部 alias FQN; - 使用 target class 声明的
module_type和exact_type校验原 module; - target class 提供
is_fusable(module)时执行额外结构兼容性检查; - 返回不可变
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
执行顺序:
- 构造全部 replacement module;
- 校验 replacement 的返回类型、train/eval 状态、
forward和 hook; - 调用可选的
replacement.make_transforms(model.config); - 为每个相对 transform 创建独立副本并设置实例作用域;
- 校验新旧参数 schema 以及 source/target 冲突;
- 确认全部 source module 仍注册在原 FQN;
- 安装 replacement;
- 将新增 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 时至少检查:
- pattern 仅描述当前 module 的相对 key;
- 原 module 消失的持久化 key 必须由 source pattern 覆盖;
- replacement 新增的持久化 key 必须由 target pattern 覆盖;
- 未参与转换的同名 parameter/buffer 仍保持原对象 identity;
- 多条规则展开后,converter source 冲突或同一 target 被重复生成时立即报错;
- 全部 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
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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