首页
/ 在 DeepSpeed Inference V2 中接入新模型:从 Parameter 到 Policy 的完整开发指南

在 DeepSpeed Inference V2 中接入新模型:从 Parameter 到 Policy 的完整开发指南

2026-09-08 21:05:24作者:殷蕙予

DeepSpeed 的推理框架中,deepspeed/inference/v2 是一套围绕 ragged batching、Tensor Parallelism 与 KV Cache 块分配全新设计的高性能推理引擎。当需要接入一个尚未被官方支持的新模型(例如新一代 Transformer 架构)时,最正统的扩展方式并不是去修改推理引擎本身,而是按照 model_implementations 子系统的既定抽象,为模型补齐“参数容器 + 模型实现 + 策略(Policy)”三个组件。本文以官方接入文档 AddingAModel.md 为主体骨架,结合仓库中真实源码(parameter_base.pylayer_container_base.pyinference_policy_base.pyinference_transformer_base.py)做纵深解读,并以内置的 llama_v2 真实实现为全程参照,帮助你掌握如何为一个类 Transformer 模型开发三件套,使其可直接被 V2 推理引擎加载并运行。

阅读并动手完成后,你将能够:为任意类 Transformer 模型编写 Parameter / LayerContainer(含 PARAM_MAPPING)、编写继承 DSTransformerModelBase 的模型实现、编写 InferenceV2Policy + ContainerMap,最终把模型接入 V2 引擎并享受其自动化的张量并行分片与参数融合优化。

一、接入新模型的三大组件与代码落位

官方文档开门见山地指出:在 DeepSpeed Inference 中接入一个新模型,需要开发三个相互关联的组件:

  1. Containers(容器):描述模型包含哪些参数;
  2. Model implementation(模型实现):描述模型应该如何被计算;
  3. Policy(策略):负责把源模型的参数映射进容器,并创建模型实现,是三者之间的“组装层”。

文档同时给出一个重要的前提假设:如果你想接入的是一个“相对传统的 Transformer 风格模型”,那么可以继承 DSTransformerModelBase,从而免费获得它提供的一整套工具方法。这里的“传统”通常意味着:embedding → N 个标准 decoder/encoder 层(self-attention + MLP + LayerNorm)→ 输出 head 的堆叠结构,而不是需要自定义稀疏注意力、非标准循环结构等特殊架构的模型。

在当前的 model_implementations 目录 中,这一“三件套”抽象已被多个真实模型采用,每个模型一个子目录,且均遵循 container.py + model.py + policy.py 的组织惯例:llama_v2mistralfalconphiphi3optqwenqwen_v2mixtralqwen_v2_moeexaone4exaone4_5 等。接入新模型时,最直接的参考样板就是 llama_v2 这套实现。

三个组件的核心类分别位于:

小提示:官方文档中的示例导入路径(如 from deepspeed.inference.module_implementations.parameter_base import ParameterBase)是早期目录规划的写法;当前仓库中这些模块已统一组织在 deepspeed.inference.v2.model_implementations 包下,本文代码示例均以仓库当前实际路径为准。

二、第一步:用 ParameterBase 定义“融合参数”

容器是“原模型参数”与“推理服务参数”之间的桥梁。但在定义容器之前,必须先理解 DeepSpeed 推理中一个最小单元——Parameter 的设计。官方文档概括了它的两个核心部分:

  • dependencies(依赖):来自源 checkpoint 的原始参数;
  • finalize 方法:当所有依赖齐备后,如何把它们转换/融合成真正的推理参数。

其动机非常典型:以 Llama 系列为例,原始 checkpoint 中 querykeyvalue 三个投影是独立存储的,但在推理时把三者沿 dim=0 融合成一个更大的 QKV 投影通常能获得更高吞吐。这个“融合”动作正是用一个自定义 Parameter 来描述的,文档给出的 UnfusedQKVParameter 示例经过导入路径修正后如下:

from deepspeed.inference.v2.model_implementations.parameter_base import ParameterBase

class UnfusedQKVParameter(ParameterBase):
    query: torch.Tensor
    key: torch.Tensor
    value: torch.Tensor

    def finalize(self) -> torch.Tensor:
        fused_param = torch.cat([self.query, self.key, self.value], dim=0)
        return self.inference_model.transform_qkv_param(fused_param)

逐段拆解这个实现:

  1. 继承 ParameterBase。基类会自动判断依赖何时被满足,并在完成后把结果写回父级 LayerContainer 的对应槽位。需要特别注意的是,ParameterBase 并非普通类:它由 ParameterMetaclass 驱动,__call__ 阶段会为每个实例初始化 dest_param=None,并把每个依赖的存储槽位初始化为 None(或空 ParametrizedList)。

  2. 类上的类型注解 = 依赖声明。元类在 __new__ 阶段解析注解,凡是注解类型是 torch.TensorParametrizedList 子类的字段,都会被识别为“依赖”,并自动替换成 property(读 getter / 写 setter)。因此类体里写 query: torch.Tensor,等价于声明“本参数依赖源 checkpoint 中一个叫 query 的张量”。由于原 Llama 模型天然是分离的 q/k/v,融合参数就依次声明三个依赖。

  3. finalize 方法。每个依赖一旦被 setter 赋值,都会调用 complete_component()completed_components 计数加一(见 parameter_base.py);当计数达到元类统计的 n_dependencies 时,finalize 被自动触发。它返回最终参数,之后该 ParameterBase 实例会被父容器的 finalization_callback 替换为真正的张量并销毁,你再也拿不到半成品。

  4. 访问依赖与调用模型工具finalize 内可以通过 self.{依赖名}(如 self.query)拿到已填充的张量;同时任何 ParameterBase 实例都持有一个 self.inference_model 弱引用——它提供了如何分片与转换参数的上下文。这里调用 self.inference_model.transform_qkv_param 正是利用了 DSTransformerBase 提供的能力:该方法会按 tp_rank/tp_size 与 head 配置对 QKV 参数做张量并行分片,随后交给底层 linear 模块实现执行量化等额外的形状转换或优化。从源码看,transform_qkv_param 先调用 shard_qkv_param(...) 完成按 head 的分片,再委托 self.qkv.transform_param(param)

2.1 内置的公共参数模板 common_parameters

由于 QKV/MLP/输出投影的融合模式在各类 Transformer 中高度重复,仓库在 common_parameters 下预置了大量可直接复用(且全部与 DSTransformerBase 兼容)的实现,接入新模型时优先查看这里:

文件 提供的典型参数模板
qkv_parameters.py FusedQKVParameter(原生已融合,直接复制)、UnfusedQKVParameter(q/k/v 分离,cat 融合)、MegatronQKVParameterGQAMegatronQKVParameter(处理 Megatron [n_heads, 3, ...] 头序布局及 GQA 分组布局转换)
attn_output_parameters.py 注意力输出投影参数(即文档示例中的 AttentionOutputParameter
mlp_parameters.py MLP 上下投影与门控参数模板
norm_parameters.py LayerNorm/RMSNorm 等归一化参数模板
embedding_parameters.pyunembed_parameters.py embedding / unembedding(LM Head)参数模板
invfreq_parameters.py RoPE inv_freq 等非张量数据的处理
moe_parameters.py MoE 路由/专家参数模板

以真实的 UnfusedQKVParameterqkv_parameters.py)为例,它与官方文档语义完全一致:torch.cat([self.q_params, self.k_params, self.v_params], dim=0) 之后接 self.inference_model.transform_qkv_param(fused_param)

2.2 变长依赖:ParametrizedList

ParameterBase 的依赖并不限于单一 torch.Tensor。某些参数的数量取决于模型配置而非架构——典型例子是 MoE 层,专家数量随模型规模变化,逐个写 expert_0expert_1… 既笨拙又不可扩展。为此 parameter_base.py 提供了 ParametrizedList

class MyParametrizedList(ParametrizedList):
    count_attr: str = "my_list_count"

其中 my_list_count 必须是推理模型实例上可访问的属性(即 self.inference_model.my_list_count),它给出该列表的长度;对列表使用整数索引 experts[8] 访问比命名成 expert_8 自然得多。该列表内部维护 set_params 计数,当所有槽位都被填充完毕才回调父参数的 complete_component()。若想省去手写子类,可以直接使用工厂函数 ParamList(attr),它会返回一个设置了指定 count_attrParametrizedList 子类。注意 ParametrizedList 不是一个普通的 listtorch.cat(param_list) 无法直接使用,需要用 torch.cat(tuple(param_list)) 包一层。

三、第二步:用 LayerContainer 组装参数并建立映射

参数单元定义好之后,下一步是把它们组合进 LayerContainer。官方文档明确了两种容器分工:

  • Transformer 容器:对应模型中的单层 Transformer 的参数,包括 FFN 全连接投影、QKV 投影、注意力输出投影、归一化等;
  • 非 Transformer 容器:存放其他一切——典型如 embedding 与 unembedding(LM Head)参数。从源码看,DSInferenceModelBase_transformer(每个 layer 一个容器的列表)和 _non_transformer(一个容器)来区分这两类参数,二者在模型实例化时都是 None,要等 Policy 调用 set_parameters(...) 之后才会被填充。

文档中简化版 Llama 的示例(仅含 QKV 与注意力输出投影),修正了语法后如下:

from deepspeed.inference.v2.model_implementations.layer_container_base import LayerContainer

class ExampleContainer(LayerContainer):
    qkvw: UnfusedQKVParameter
    attn_o: AttentionOutputParameter

    PARAM_MAPPING = {
        "self_attn.q_proj.weight": "qkvw.query",
        "self_attn.k_proj.weight": "qkvw.key",
        "self_attn.v_proj.weight": "qkvw.value",
        "self_attn.o_proj.weight": "attn_o.params",
    }

同样有两个关键要素:

  1. 参数类型注解。每个注解对应模型实现中可以使用的“参数组”。在模型实现的 forward 中,直接写 container.qkvw 就能拿到已经完成融合、分片与变换的 QKV 参数(它实际上已被替换为 InferenceParameter 张量)。之所以能做到这一点,是因为 LayerContainer 同样由元类驱动(LayerMetaclass):它在 __new__ 阶段收集本类与所有基类的注解并合并(因此可以通过继承复用公共参数组),为每个 ParameterBase 注解生成实例与统一的 finalization_callback

  2. PARAM_MAPPING 字典。这是“源 checkpoint 参数名 → 容器内依赖”的显式路由表。它会被 Policy 利用,在加载 checkpoint 时自动填充依赖,因此文档强调:它是显式映射,写错、漏写都会被直接拦截。事实上 LayerMetaclass 在构建阶段会做大量静态校验(见 layer_container_base.py):

    • 映射目标必须写成 参数名.依赖名,且 参数名 必须是本容器注解中真实存在的参数,依赖名必须真实存在于该参数的注解中;
    • 同一个依赖不允许被多条映射规则命中(重复映射直接抛 ValueError);
    • 容器中所有 ParameterBase 的所有依赖必须被映射覆盖完全,否则会以“以下依赖未被映射”的明确错误拒绝;
    • 若一个依赖目标是 ParametrizedList,则源名必须带且只能带一个通配符 *(如 "model.layers.*...."),元类会据此生成把源名中的数字索引解析到列表下标的路由 helper;禁止把 ParametrizedList 与普通 Tensor 混在同一个映射规则里。

在 checkpoint 加载期间,容器会通过 set_dependency(dep_name, dep_value)layer_container_base.py)接收去除了前缀的依赖名:先尝试精确匹配 PARAM_MAPPING,再尝试通配符(* 会被替换为正则 .*)与 plist_helpers 的列表索引解析。此外容器还提供 direct_injection(name, tensor) 用于直接注入张量,以及两个非常有用的只读状态属性:is_populated(所有参数都已被 checkpoint 引擎填充)与 is_initialized(在填充基础上进一步要求全部落在正确的加速设备上,参数类型必须是 InferenceParameter 或显式 None)。

文档建议:当 Transformer 容器与非 Transformer 容器都写好之后,就可以进入模型实现环节了。

四、第三步:编写继承 DSTransformerModelBase 的模型实现

DSTransformerModelBaseinference_transformer_base.py)承担了绝大部分“分片与变换参数”的机械工作:在其 __init__ 中,它按固定顺序调用 make_norm_layer()make_qkv_layer()make_attn_layer()make_attn_out_layer()make_mlp_1_layer()make_mlp_2_layer()make_embedding_layer()make_unembedding_layer(),为模型组装出 embedding、QKV、self-attention、注意力输出投影、两层 MLP、归一化与 LM Head 的整套模块,并借助 modules.heuristics.instantiate_* 依据引擎配置(ragged state manager 的 max_ragged_batch_size 等)选择具体 kernel 实现。

即便如此,接入者仍需要亲自完成官方文档点名的四个关键任务

任务 1:根据模型配置定义抽象属性

DSTransformerModelBase 用大量 @property @abstractmethod 把模型“尺寸”与“结构”参数化,接入者必须全部实现。从源码归纳,至少包括:

  • 尺寸类num_layers(层数)、model_dim(embedding 与残差维度)、vocab_size(含 padding 的词典大小)、head_size(每个注意力头维度)、n_heads(query 头数)、intermediate_dim(未分片的中间投影维度;对门控激活,指第二个 MLP 层的输入维度);
  • 结构类activation_dtypemlp_activation_fn(MLP 激活函数,决定是否为门控)、norm_typepositional_embedding_type(位置编码类型)与 positional_embedding_config(通常为 RoPE 的 RotateHalfConfig)。

基类还基于这些属性派生了便捷工具:n_heads_q_local / n_heads_kv_local 给出按 tp_rank/tp_size 分片后的本地头数;gated_mlp 判断是否使用门控激活。需要特别留意 n_heads_kv:默认实现采用 MHA 形式(return self.n_heads),GQA 或 MQA 模型必须重写该属性

任务 2:配置 embedding / unembedding 模块并实现其 forward

  • make_embedding_layer() 默认只做 dtype 转换,不沿 channel 维度分片(源码注释说明在支持非连续 all-gather 之前不会分片 embedding 参数);
  • make_unembedding_layer() 假设 LM Head 之前存在一次归一化,并对 vocab 维做分片(sharded_unembed_dim),tp_size > 1 时还会预分配通信用的 logits 缓冲。

如果模型不符合这些默认假设(例如无 pre-norm、embedding 需要特殊处理),就应重写对应 make_* 方法,并配套编写 embedding 与 unembedding 的 forward 计算。

任务 3:配置注意力与 KV Cache 行为

基类为自回归密集注意力模式提供了一整套可覆盖的钩子:

  • make_attn_layer()softmax_scale = 1.0 / head_size**0.5、本地头数与位置编码信息构建 DSSelfAttentionConfig
  • kv_cache_config() 返回 KVCacheConfig,其中 cache_shape = (num_layers, n_heads_kv_local, head_size)max_blocks_per_allocation_groupmax_sequence_lengthkv_block_size 推算;
  • get_kv_requirements(...) / maybe_allocate_kv(...) 负责估算并在需要时向 state_manager 申请 KV 块;
  • prepare_batch(...) 在每次 forward 前构造 attention 相关的批元数据(如 build_atoms)。

若模型的注意力模式不是标准的自回归密集注意力(例如有滑动窗口、稀疏模式),文档与源码均明确要求重写这些默认实现。

任务 4:编写 Transformer 层的 forward

最后是模型实现真正“计算如何发生”的部分:以容器为输入,按 归一化 → 注意力(自注意力 + 输出投影 + 残差)→ MLP(两层投影 + 残差)→ 输出 LM 计算 的顺序编写单层 forward,以及全模型的逐层循环逻辑。注意 inference_model_base.pyforward 的接口约定:它需要能被构图(graphable),因此不应依赖 Python 控制流编写逐 token 的逻辑。

另外,仓库还在同一文件中提供了 DSMoETransformerModelBase(MoE 版基类),要求补充 n_expertsn_top_knormalize_expert_scores 属性,并用 make_moe_layer() 替代普通 MLP,同时区分 transform_moe_mlp_1_paramtransform_moe_mlp_2_param(因为同一 DSModule 同时持有两块专家参数时,无法仅凭形状推断应做哪种变换)。mixtralqwen_v2_moe 目录就是这一基类的最佳实践。

五、第四步:编写 Policy 与 ContainerMap

官方文档对 Policy 的定位是“组合层”:InferenceV2Policy 是直接传给推理引擎的对象,负责把模型实现与容器组合成端到端解决方案。仓库中该抽象定义在 inference_policy_base.py,其构造函数要求 checkpoint_engineinf_checkpoint_path 二选一(同时为 None 或同时给出都会抛 ValueError),分别对应两种建模型机制:

  1. checkpoint_engine:通用机制,checkpoint 引擎逐个遍历模型参数交给 Policy,再由模型实现完成分片/变换;
  2. inf_checkpoint_path:重新加载此前由 DeepSpeed 序列化保存的推理模型(这种 checkpoint 不应跨模型后端配置混用,源码以 TODO 注明该限制尚待代码强制)。

Policy 需要实现两个抽象方法:

5.1 instantiate_model:创建模型实例

第一个抽象方法是创建前面定义好的模型。文档指出一般情况下只需直接调用模型构造函数,并把引擎配置、张量并行通信对象、自定义模型配置三个参数传进去即可。这与仓库中 llama_v2 的实现完全吻合(policy.py):

def instantiate_model(self, engine_config: RaggedInferenceEngineConfig, mp_group: Any) -> Llama2InferenceModel:
    return Llama2InferenceModel(config=self._model_config, engine_config=engine_config, base_mp_group=mp_group)

5.2 build_container_map:建立 checkpoint 前缀 → 容器的路由

第二个抽象方法是定义 checkpoint 参数如何映射到每个容器。上一节提过,LayerContainer 自己就能处理“checkpoint 参数 → 容器内依赖”的内部路由(PARAM_MAPPING),但为了找到“该参数到底属于哪一层、哪个容器”,还需要 ContainerMap 这一层抽象。

ContainerMapinference_policy_base.py)通过把“checkpoint 前缀字符串”分类到“它所对应的容器类型”来完成映射,核心 API 有三个:

  • set_transformer_params(prefixes, containers):注册 transformer 前缀与其对应的每层一个的容器列表(入参必须是一个 list,容器数与层数一一对应);
  • set_non_transformer_params(container):注册唯一的非 Transformer 容器;
  • set_unmapped_params(prefixes):登记“剩余”的、应当被忽略的前缀(比如属于 RoPE 缓存等运行时数据、不属于权重的张量)。

在加载阶段,ContainerMap.map_param(name, parameter) 会对每个参数名依次做三件事:先检查是否命中“忽略前缀”;再检查是否命中 transformer 前缀——命中则剥离前缀并解析层号,把剩余名字路由到第 layer_idx 个 transformer 容器的 set_dependency;否则尝试交给 _non_transformer_params,若连它也抛 ValueError,则包装成信息更明确的报错:“找不到该参数对应的容器,请复查 Containers/ContainerMap”。完成全部填充后调用 validate(),逐一校验非 Transformer 容器与每个 transformer 容器都已 is_initialized

最简洁的 build_container_map 通常是遍历某个 PyTorch 模型的 named_parameters(),或遍历一个 checkpoint 的 state dict 后依前缀归类。官方文档示例中 set_transformer_params("model.layers", transformer_containers) 正是把 model.layers 前缀下的参数路由给每一层容器。下面是 llama_v2 的真实实现,可作为模板(policy.py):

def build_container_map(self) -> ContainerMap:
    map = ContainerMap()

    transformer_containers = [Llama2TransformerContainer(self.model) for _ in range(self.model.num_layers)]
    map.set_transformer_params(['model.layers'], transformer_containers)

    map.set_non_transformer_params(Llama2NonTransformerContainer(self.model))

    map.set_unmapped_params(
        [f'model.layers.{i}.self_attn.rotary_emb.inv_freq' for i in range(self.model.num_layers)])

    return map

5.3 引擎如何驱动 Policy

当引擎需要建模型时,会调用 InferenceV2Policy.build_model(engine_config, mp_group)inference_policy_base.py),其内部流程为:

  1. self.model = self.instantiate_model(engine_config, mp_group) 创建模型实现;
  2. self.populate_model_parameters()build_container_map(),然后按上文两种机制之一填充参数:若是 checkpoint 引擎,则遍历 checkpoint_engine.parameters() 逐个 map_param,随后用 flat_model_helpers.pyflatten_inference_model 把容器参数拍平为一个连续 buffer 与其元数据;若是序列化路径,则按 tp_rank/tp_size 定位 make_param_filename / make_metadata_filename 生成的文件并 restore_inference_model
  3. container_map.validate() 校验完整性;
  4. self.model.set_parameters(...) 把 transformer / non-transformer 容器与拍平后的参数 buffer 交给模型。

另外值得留意的是,PolicyMeta 元类会自动把每个非基类的 Policy 子类注册进模块级 POLICIES 字典——引擎侧据此按名字发现可用的模型策略。

六、推荐阅读路径与调试建议

动手实现时,建议按以下顺序在仓库中“对照着抄”:

  1. 先读一套最完整的官方样例llama_v2/container.py(看 Llama2TransformerContainer / Llama2NonTransformerContainer 如何写注解与 PARAM_MAPPING)、llama_v2/model.py(看如何实现抽象属性与逐层 forward)、llama_v2/policy.py(看三行核心方法);
  2. 相似架构优先拷贝参数模板:GQA 类(Llama 2/3、Qwen2、Mistral 等)从 common_parameters/qkv_parameters.py 选取 UnfusedQKVParameter / GQAMegatronQKVParameter;Megatron 风格布局参考 MegatronQKVParameter;MoE 架构参考 mixtralqwen_v2_moe
  3. 写完后自查容器完整性:由于 LayerMetaclass 在类定义期就会强制“依赖全覆盖 + 无重复映射 + 目标存在”,绝大多数映射笔误会在实例化容器的那一刻直接暴露;加载期若报“Cannot find container for {name}”,说明 PARAM_MAPPINGContainerMap 的源名前缀不匹配——对照上述 llama 参考,仔细核对 model.layers.{i}.* 之类的层前缀即可;
  4. 设备相关报错查 is_initialized:若容器校验提示参数不在正确设备,检查你的 finalize 链路里参数是否都经过 transform_*_param(如 transform_qkv_paramtransform_mlp_1_param)并落在 InferenceParameter 上。

接入文档本身保存在 AddingAModel.md,与各模型实现源码同目录,非常适合作为“写代码时的案头手册”。整体而言,这套“Parameter 声明融合 → LayerContainer 路由参数 → DSTransformerModelBase 计算 → Policy 组装”的分层设计,把「checkpoint 参数布局的差异」和「计算图/分片/量化等推理细节」彻底解耦——新增模型时大部分工作集中在描述参数映射关系上,而分片、融合、KV 缓存管理等繁琐部分则由框架基类自动接管。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
docsdocs
暂无描述
Markdown
899
5.83 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
925
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.84 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
601
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.03 K
525
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
395