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 传递路径。
模型高性能实现替换:
weights_mapping接入1. 当前代码基线
新代码已经实现 PR #1156 的后构造 module replacement,真实调用链是:
当前
apply_module_replacements()接口直接返回 model:apply_module_replacements( model: nn.Module, plan: ModuleReplacementPlan, *, context: Mapping[str, Any] | None = None, ) -> nn.Modulereplacement target 使用
@module_replacement声明,接收module、module_fqn、只读context,当前必须构造并返回nn.Module。apply_module_replacements()会原地更新父模块的_modules,最后返回同一个根 model。当前
_validate_replacement()只允许结构保持替换: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.ModuleFQN。因此 target 也必须构造nn.Module,以承接原模块的 Parameter、buffer、training state、hooks、state_dict和forward协议。npu_rms_norm是计算函数,不是 model tree 节点:它没有 module FQN,也不持有 Parameter,当前 executor 无法用 module FQN 直接替换它。正确路径是:直接 target 到
npu_rms_norm需要 codegen 暴露函数调用点,并由另一套 call-site replacement 机制处理;它不属于当前 module replacement 协议。module_type支持一个nn.Moduleclass 或 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。输入与 TransformersModuleFusionSpec.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,与 TransformersModuleFusionSpec.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;module_type和exact_type校验原 module;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执行顺序:
forward和 hook;replacement.make_transforms(model.config);实例作用域使用
WeightTransform已有字段:共享 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_typemapping。直接把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)state_dictkey 全部不变forward兼容性、hook 限制和 transform 完整性有 transforms 时至少检查:
make_transforms()的返回值必须是list[WeightTransform],具体 transform 的构造和反向转换语义沿用 Transformers 定义。6. 实现改动位置
components/model_transform/replacement.pyis_fusable;apply 阶段构造 replacement、调用可选make_transforms(model.config)、校验 scoped transforms,并注册到 Transformers root model mappingtrainer/config.pymodule_type/exact_type;从 target decorator 元数据构造 specis_fusable、无 transforms、融合 transforms、多个 FQN、alias、冲突、原子失败及get_model_conversion_mapping(model)可见性7. Model transform 验收
验收覆盖 replacement class 构造、
make_transforms(model.config)调用、FQN scope、schema 覆盖、冲突检查、原子安装,以及新增 transforms 能由 Transformers 原生get_model_conversion_mapping(model)读取;checkpoint loader 保持现状。