在 DeepSpeed Inference V2 中接入新模型:从 Parameter 到 Policy 的完整开发指南
在 DeepSpeed 的推理框架中,deepspeed/inference/v2 是一套围绕 ragged batching、Tensor Parallelism 与 KV Cache 块分配全新设计的高性能推理引擎。当需要接入一个尚未被官方支持的新模型(例如新一代 Transformer 架构)时,最正统的扩展方式并不是去修改推理引擎本身,而是按照 model_implementations 子系统的既定抽象,为模型补齐“参数容器 + 模型实现 + 策略(Policy)”三个组件。本文以官方接入文档 AddingAModel.md 为主体骨架,结合仓库中真实源码(parameter_base.py、layer_container_base.py、inference_policy_base.py 与 inference_transformer_base.py)做纵深解读,并以内置的 llama_v2 真实实现为全程参照,帮助你掌握如何为一个类 Transformer 模型开发三件套,使其可直接被 V2 推理引擎加载并运行。
阅读并动手完成后,你将能够:为任意类 Transformer 模型编写 Parameter / LayerContainer(含 PARAM_MAPPING)、编写继承 DSTransformerModelBase 的模型实现、编写 InferenceV2Policy + ContainerMap,最终把模型接入 V2 引擎并享受其自动化的张量并行分片与参数融合优化。
一、接入新模型的三大组件与代码落位
官方文档开门见山地指出:在 DeepSpeed Inference 中接入一个新模型,需要开发三个相互关联的组件:
- Containers(容器):描述模型包含哪些参数;
- Model implementation(模型实现):描述模型应该如何被计算;
- Policy(策略):负责把源模型的参数映射进容器,并创建模型实现,是三者之间的“组装层”。
文档同时给出一个重要的前提假设:如果你想接入的是一个“相对传统的 Transformer 风格模型”,那么可以继承 DSTransformerModelBase,从而免费获得它提供的一整套工具方法。这里的“传统”通常意味着:embedding → N 个标准 decoder/encoder 层(self-attention + MLP + LayerNorm)→ 输出 head 的堆叠结构,而不是需要自定义稀疏注意力、非标准循环结构等特殊架构的模型。
在当前的 model_implementations 目录 中,这一“三件套”抽象已被多个真实模型采用,每个模型一个子目录,且均遵循 container.py + model.py + policy.py 的组织惯例:llama_v2、mistral、falcon、phi、phi3、opt、qwen、qwen_v2、mixtral、qwen_v2_moe、exaone4、exaone4_5 等。接入新模型时,最直接的参考样板就是 llama_v2 这套实现。
三个组件的核心类分别位于:
- 参数单元抽象:parameter_base.py 中的
ParameterBase; - 参数容器抽象:layer_container_base.py 中的
LayerContainer; - 策略与映射抽象:inference_policy_base.py 中的
InferenceV2Policy与ContainerMap; - 模型基类:inference_model_base.py 的
DSInferenceModelBase,以及 inference_transformer_base.py 的DSTransformerModelBase。
小提示:官方文档中的示例导入路径(如
from deepspeed.inference.module_implementations.parameter_base import ParameterBase)是早期目录规划的写法;当前仓库中这些模块已统一组织在deepspeed.inference.v2.model_implementations包下,本文代码示例均以仓库当前实际路径为准。
二、第一步:用 ParameterBase 定义“融合参数”
容器是“原模型参数”与“推理服务参数”之间的桥梁。但在定义容器之前,必须先理解 DeepSpeed 推理中一个最小单元——Parameter 的设计。官方文档概括了它的两个核心部分:
- dependencies(依赖):来自源 checkpoint 的原始参数;
- finalize 方法:当所有依赖齐备后,如何把它们转换/融合成真正的推理参数。
其动机非常典型:以 Llama 系列为例,原始 checkpoint 中 query、key、value 三个投影是独立存储的,但在推理时把三者沿 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)
逐段拆解这个实现:
-
继承
ParameterBase。基类会自动判断依赖何时被满足,并在完成后把结果写回父级LayerContainer的对应槽位。需要特别注意的是,ParameterBase并非普通类:它由ParameterMetaclass驱动,__call__阶段会为每个实例初始化dest_param=None,并把每个依赖的存储槽位初始化为None(或空ParametrizedList)。 -
类上的类型注解 = 依赖声明。元类在
__new__阶段解析注解,凡是注解类型是torch.Tensor或ParametrizedList子类的字段,都会被识别为“依赖”,并自动替换成 property(读 getter / 写 setter)。因此类体里写query: torch.Tensor,等价于声明“本参数依赖源 checkpoint 中一个叫query的张量”。由于原 Llama 模型天然是分离的 q/k/v,融合参数就依次声明三个依赖。 -
finalize方法。每个依赖一旦被 setter 赋值,都会调用complete_component()把completed_components计数加一(见 parameter_base.py);当计数达到元类统计的n_dependencies时,finalize被自动触发。它返回最终参数,之后该ParameterBase实例会被父容器的finalization_callback替换为真正的张量并销毁,你再也拿不到半成品。 -
访问依赖与调用模型工具。
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 融合)、MegatronQKVParameter 与 GQAMegatronQKVParameter(处理 Megatron [n_heads, 3, ...] 头序布局及 GQA 分组布局转换) |
| attn_output_parameters.py | 注意力输出投影参数(即文档示例中的 AttentionOutputParameter) |
| mlp_parameters.py | MLP 上下投影与门控参数模板 |
| norm_parameters.py | LayerNorm/RMSNorm 等归一化参数模板 |
| embedding_parameters.py、unembed_parameters.py | embedding / unembedding(LM Head)参数模板 |
| invfreq_parameters.py | RoPE inv_freq 等非张量数据的处理 |
| moe_parameters.py | MoE 路由/专家参数模板 |
以真实的 UnfusedQKVParameter(qkv_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_0、expert_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_attr 的 ParametrizedList 子类。注意 ParametrizedList 不是一个普通的 list:torch.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",
}
同样有两个关键要素:
-
参数类型注解。每个注解对应模型实现中可以使用的“参数组”。在模型实现的 forward 中,直接写
container.qkvw就能拿到已经完成融合、分片与变换的 QKV 参数(它实际上已被替换为InferenceParameter张量)。之所以能做到这一点,是因为LayerContainer同样由元类驱动(LayerMetaclass):它在__new__阶段收集本类与所有基类的注解并合并(因此可以通过继承复用公共参数组),为每个ParameterBase注解生成实例与统一的finalization_callback。 -
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 的模型实现
DSTransformerModelBase(inference_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_dtype、mlp_activation_fn(MLP 激活函数,决定是否为门控)、norm_type、positional_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_group由max_sequence_length与kv_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.py 中 forward 的接口约定:它需要能被构图(graphable),因此不应依赖 Python 控制流编写逐 token 的逻辑。
另外,仓库还在同一文件中提供了 DSMoETransformerModelBase(MoE 版基类),要求补充 n_experts、n_top_k、normalize_expert_scores 属性,并用 make_moe_layer() 替代普通 MLP,同时区分 transform_moe_mlp_1_param 与 transform_moe_mlp_2_param(因为同一 DSModule 同时持有两块专家参数时,无法仅凭形状推断应做哪种变换)。mixtral、qwen_v2_moe 目录就是这一基类的最佳实践。
五、第四步:编写 Policy 与 ContainerMap
官方文档对 Policy 的定位是“组合层”:InferenceV2Policy 是直接传给推理引擎的对象,负责把模型实现与容器组合成端到端解决方案。仓库中该抽象定义在 inference_policy_base.py,其构造函数要求 checkpoint_engine 与 inf_checkpoint_path 二选一(同时为 None 或同时给出都会抛 ValueError),分别对应两种建模型机制:
checkpoint_engine:通用机制,checkpoint 引擎逐个遍历模型参数交给 Policy,再由模型实现完成分片/变换;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 这一层抽象。
ContainerMap(inference_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),其内部流程为:
self.model = self.instantiate_model(engine_config, mp_group)创建模型实现;self.populate_model_parameters()调build_container_map(),然后按上文两种机制之一填充参数:若是 checkpoint 引擎,则遍历checkpoint_engine.parameters()逐个map_param,随后用 flat_model_helpers.py 的flatten_inference_model把容器参数拍平为一个连续 buffer 与其元数据;若是序列化路径,则按tp_rank/tp_size定位make_param_filename/make_metadata_filename生成的文件并restore_inference_model;container_map.validate()校验完整性;self.model.set_parameters(...)把 transformer / non-transformer 容器与拍平后的参数 buffer 交给模型。
另外值得留意的是,PolicyMeta 元类会自动把每个非基类的 Policy 子类注册进模块级 POLICIES 字典——引擎侧据此按名字发现可用的模型策略。
六、推荐阅读路径与调试建议
动手实现时,建议按以下顺序在仓库中“对照着抄”:
- 先读一套最完整的官方样例:llama_v2/container.py(看
Llama2TransformerContainer/Llama2NonTransformerContainer如何写注解与PARAM_MAPPING)、llama_v2/model.py(看如何实现抽象属性与逐层 forward)、llama_v2/policy.py(看三行核心方法); - 相似架构优先拷贝参数模板:GQA 类(Llama 2/3、Qwen2、Mistral 等)从 common_parameters/qkv_parameters.py 选取
UnfusedQKVParameter/GQAMegatronQKVParameter;Megatron 风格布局参考MegatronQKVParameter;MoE 架构参考 mixtral 与 qwen_v2_moe; - 写完后自查容器完整性:由于
LayerMetaclass在类定义期就会强制“依赖全覆盖 + 无重复映射 + 目标存在”,绝大多数映射笔误会在实例化容器的那一刻直接暴露;加载期若报“Cannot find container for {name}”,说明PARAM_MAPPING与ContainerMap的源名前缀不匹配——对照上述 llama 参考,仔细核对model.layers.{i}.*之类的层前缀即可; - 设备相关报错查
is_initialized:若容器校验提示参数不在正确设备,检查你的finalize链路里参数是否都经过transform_*_param(如transform_qkv_param、transform_mlp_1_param)并落在InferenceParameter上。
接入文档本身保存在 AddingAModel.md,与各模型实现源码同目录,非常适合作为“写代码时的案头手册”。整体而言,这套“Parameter 声明融合 → LayerContainer 路由参数 → DSTransformerModelBase 计算 → Policy 组装”的分层设计,把「checkpoint 参数布局的差异」和「计算图/分片/量化等推理细节」彻底解耦——新增模型时大部分工作集中在描述参数映射关系上,而分片、融合、KV 缓存管理等繁琐部分则由框架基类自动接管。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0631
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
video-shotcraftAI宣传片skill,使用 Remotion 制作电影级产品视频:提供106 张镜头配方卡和可复用的视频魔板。适用于 Claude Code 与 Codex以及所有其他智能体Markdown00
HivisionIDPhotos⚡️HivisionIDPhotos: a lightweight and efficient AI ID photos tools. 一个轻量级的AI证件照制作算法。Python09
DragonOSDragonOS is an operating system developed from scratch using Rust, with Linux compatibility. It is designed for **Serverless** scenarios. 使用Rust从0自研内核,具有Linux兼容性的操作系统,面向云计算Serverless场景而设计。Rust00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00