执行摘要
- 一句话:Transformers 后端输入嵌入处理全面泛化,新增 replace_embedding_class
- 推荐动作:值得精读(尤其是
replace_embedding_class 的实现),它展示了用 MRO 重排和 __class__ 赋值优雅处理第三方子类替换的设计模式,且与 LoRA、量化、TP 的交互均有考虑。建议关注后续真实模型端到端验证是否补全。
功能与动机
PR body 明确指出旧实现的问题:"Before this PR we replaced the entire module returned by get_input_embeddings with VocabParallelEmbedding and had a special case for scaled input embeddings. This does not generalise well, particularly if input embeddings perform additional operations that we have not accounted for." 意味着仅支持裸 nn.Embedding 或单一缩放嵌入,无法处理嵌入内执行额外操作(如组合包装、额外变换)的模型。
实现拆解
- 在
vllm/model_executor/models/transformers/utils.py 中新增 replace_embedding_class 主函数,逻辑分三层:
- 若嵌入模块不是
nn.Embedding 实例,则递归查找其内部恰好一个 nn.Embedding 子模块并替换,通过新引入的 attrsetter 支持嵌套属性设置;若组合数量不为 1 则抛出 ValueError。
- 若嵌入是裸
nn.Embedding(类型精确匹配),直接构造 VocabParallelEmbedding 实例替换。
- 若嵌入继承自
nn.Embedding,则通过 _rebase_on_vocab_parallel 动态创建新类,将 VocabParallelEmbedding 插入 MRO 前方,并借助 _UninitializedEmbedding 使 nn.Embedding.__init__ 失效,防止重复初始化,最后原地修改 embedding.__class__ 并手动调用 VocabParallelEmbedding.__init__。
- 在
base.py 中移除 ScaledVocabParallelEmbedding 类和原先基于 getattr_iter 的特判逻辑,改为直接调用 replace_embedding_class,精简约 39 行。
causal.py 中 tie_word_embeddings 分支不再直接把包含包装器的整个输入嵌入传给 tie_weights,而是遍历 get_input_embeddings().modules() 找到唯一 VocabParallelEmbedding 实例进行权重绑定。
vllm/model_executor/layers/vocab_parallel_embedding.py 将 is_embedding_layer 判断从 type(self) is VocabParallelEmbedding 改为 not isinstance(self, ParallelLMHead),兼容集成后派生类仍被识别为嵌入层。
vllm/lora/layers/vocal_parallel_embedding.py 的 can_replace_layer 由精确类型比较改为 isinstance 检查,并排除 ParallelLMHead,支持 Transformers 后端派生的嵌入类被 LoRA 替换。
- 测试配套:
tests/models/transformers/test_backend.py 新增 176 行单元测试,覆盖裸嵌入、继承、组合、嵌套、歧义组合及 lm_head 绑定等场景。
关键文件:
vllm/model_executor/models/transformers/utils.py(模块 模型执行器;类别 source;类型 core-logic;符号 attrsetter, setter, _UninitializedEmbedding, _VocabParallelEmbeddingBase): 核心改动文件,新增 replace_embedding_class 主函数、attrsetter、_UninitializedEmbedding、_VocabParallelEmbeddingBase 和 _rebase_on_vocab_parallel,定义了三类嵌入的替换策略。
vllm/model_executor/models/transformers/base.py(模块 模型执行器;类别 source;类型 core-logic;符号 ScaledVocabParallelEmbedding, init): 移除 ScaledVocabParallelEmbedding 特判,改用 replace_embedding_class 统一处理输入嵌入,是行为变化的主入口。
tests/models/transformers/test_backend.py(模块 测试套件;类别 test;类型 test-coverage;符号 ScaledWordEmbedding, ComposedWordEmbedding, replace, assert_scaled): 新增 176 行单元测试,覆盖 replace_embedding_class 的裸、继承、组合、嵌套、歧义抛错和 lm_head 绑定场景,是该 PR 正确性的关键保障。
vllm/model_executor/models/transformers/causal.py(模块 模型执行器;类别 source;类型 core-logic): tie_word_embeddings 绑定逻辑从直接绑定整个输入嵌入改为遍历查找 VocabParallelEmbedding,是修复 Codex 指出的 P1 问题的关键。
vllm/lora/layers/vocal_parallel_embedding.py(模块 LoRA 层;类别 source;类型 dependency-wiring): LoRA 层可替换检测从精确类型变为 isinstance,并排除 ParallelLMHead,确保 Transformers 后端 rebase 产生的子类能被 LoRA 正确识别。
vllm/model_executor/layers/vocab_parallel_embedding.py(模块 模型执行器;类别 source;类型 core-logic): is_embedding_layer 判断从精确类型改为排除 ParallelLMHead,使得 Transformers 后端动态子类仍走嵌入量化路径。
关键符号:replace_embedding_class, attrsetter, _rebase_on_vocab_parallel, CausalMixin.init, VocabParallelEmbeddingWithLoRA.can_replace_layer
关键源码片段
vllm/model_executor/models/transformers/utils.py
核心改动文件,新增 replace_embedding_class 主函数、attrsetter、_UninitializedEmbedding、_VocabParallelEmbeddingBase 和 _rebase_on_vocab_parallel,定义了三类嵌入的替换策略。
# vllm/model_executor/models/transformers/utils.py
def attrsetter(attr: str) -> Callable[[object, object], None]:
"""设置可能嵌套的属性,作为 attrgetter 的逆操作。"""
parent, _, name = attr.rpartition(".")
def setter(obj: object, value: object):
# 先解析父对象,再设置最末层属性,支持 "inner.embed" 这类路径
attr_parent = attrgetter(parent)(obj) if parent else obj
setattr(attr_parent, name, value)
return setter
class _UninitializedEmbedding(nn.Embedding):
"""让 `nn.Embedding.__init__` 失效。
目的是当 `VocabParallelEmbedding.__init__` 调用 `super().__init__` 时,
不会真的执行 `nn.Embedding.__init__`,避免重复初始化权重。
"""
def __init__(self):
pass
class _VocabParallelEmbeddingBase(VocabParallelEmbedding, _UninitializedEmbedding):
"""通过 MRO 顺序让 `VocabParallelEmbedding.forward` 优先于 `nn.Embedding.forward`。
这样 `nn.Embedding` 子类中的 `super().forward(...)` 会先到达 vLLM 的嵌入实现,
从而保留子类自身的缩放或其他额外行为。
"""
@lru_cache
def _rebase_on_vocab_parallel(cls: type[nn.Embedding]) -> type[VocabParallelEmbedding]:
"""为 `cls` 动态创建新类,使其继承 `_VocabParallelEmbeddingBase`。
缓存确保同一个 `cls` 始终映射到同一个动态类,避免重复创建。
返回的类会赋值给实例的 `__class__`。
"""
return type(cls.__name__, (cls, _VocabParallelEmbeddingBase), {})
def replace_embedding_class(
embedding: nn.Module,
quant_config: "QuantizationConfig | None" = None,
*,
prefix: str = "",
) -> nn.Module:
"""将 `embedding` 中的 `nn.Embedding` 替换为 `VocabParallelEmbedding`。"""
# 如果 embedding 本身不是 nn.Embedding,说明它组合(composition)了一个,
# 需要递归找到并替换内部那个 nn.Embedding
if not isinstance(embedding, nn.Embedding):
composed = [
(name, module)
for name, module in embedding.named_modules()
if isinstance(module, nn.Embedding)
]
# 只允许恰好一个 nn.Embedding,否则权重归属不明确,直接报错
if len(composed) != 1:
raise ValueError(
f"Expected {type(embedding).__name__} to be an `nn.Embedding` or to "
f"compose exactly one, but found {len(composed)}."
)
name, module = composed[0]
new_embedding = replace_embedding_class(
module, quant_config, prefix=maybe_prefix(prefix, name)
)
# 用 attrsetter 支持嵌套名称(如 "inner.embed"),并原地替换
attrsetter(name)(embedding, new_embedding)
return embedding
# 构造 VocabParallelEmbedding 所需的参数,形状和 dtype 均取自原模块
kwargs = dict(
num_embeddings=embedding.num_embeddings,
embedding_dim=embedding.embedding_dim,
params_dtype=embedding.weight.dtype,
quant_config=quant_config,
prefix=prefix,
)
# 裸 nn.Embedding:直接整体替换成新的 VocabParallelEmbedding
if type(embedding) is nn.Embedding:
return VocabParallelEmbedding(**kwargs)
# 继承 nn.Embedding:保留子类额外状态与 forward 行为,原地 rebase
embedding.__class__ = _rebase_on_vocab_parallel(type(embedding))
# 手动调用父类初始化,跳过 nn.Embedding.__init__(已被 _UninitializedEmbedding 屏蔽)
VocabParallelEmbedding.__init__(embedding, **kwargs)
return embedding
vllm/model_executor/models/transformers/base.py
移除 ScaledVocabParallelEmbedding 特判,改用 replace_embedding_class 统一处理输入嵌入,是行为变化的主入口。
# vllm/model_executor/models/transformers/base.py
# Base.__init__ 中输入嵌入替换的简化逻辑
# Input embeddings
input_embeddings = self.model.get_input_embeddings()
if not isinstance(input_embeddings, PPMissingLayer):
# 统一走 replace_embedding_class:
# - 裸 nn.Embedding 直接替换为 VocabParallelEmbedding
# - 继承 / 组合的嵌入保留额外行为,仅替换内部或 rebase 类
# 不再需要手动从 config 读 vocab_size / hidden_size,
# 形状信息直接取自原嵌入模块自身,避免 config 与模型不一致
self.model.set_input_embeddings(
replace_embedding_class(input_embeddings, self.quant_config)
)
tests/models/transformers/test_backend.py
新增 176 行单元测试,覆盖 replace_embedding_class 的裸、继承、组合、嵌套、歧义抛错和 lm_head 绑定场景,是该 PR 正确性的关键保障。
# tests/models/transformers/test_backend.py
# 核心测试替身与断言辅助
VOCAB_SIZE = 64
HIDDEN_SIZE = 8
EMBED_SCALE = 3.0
class ScaledWordEmbedding(nn.Embedding):
"""模拟 Transformers 中带缩放的嵌入子类(含额外 buffer)。"""
def __init__(
self, num_embeddings, embedding_dim, padding_idx=None, embed_scale=1.0
):
super().__init__(num_embeddings, embedding_dim, padding_idx)
self.scalar_embed_scale = embed_scale
# 持久化禁用,避免干扰 state_dict 比较
self.register_buffer("embed_scale", torch.tensor(embed_scale), persistent=False)
def forward(self, input_ids):
# 子类 forward 中调用 super().forward,期望最终命中 VocabParallelEmbedding
return super().forward(input_ids) * self.embed_scale.to(self.weight.dtype)
class ComposedWordEmbedding(nn.Module):
"""组合型嵌入:不继承 nn.Embedding,而是包装一个 .embed 子模块。"""
def __init__(self, num_embeddings, embedding_dim, embed_scale=1.0):
super().__init__()
self.embed = nn.Embedding(num_embeddings, embedding_dim)
self.embed_scale = embed_scale
def forward(self, input_ids):
return self.embed(input_ids) * self.embed_scale
def assert_scaled(vpe, module, embedding=None):
"""断言 module 的输出等于内部嵌入的无缩放输出乘以 EMBED_SCALE。"""
input_ids = torch.arange(VOCAB_SIZE)
# 用未被放大的 VocabParallelEmbedding 前向作为基准
unscaled = vpe.forward(embedding if embedding is not None else module, input_ids)
torch.testing.assert_close(module(input_ids), unscaled * EMBED_SCALE)
评论区精华
- 组合嵌入的 LM Head 绑定问题(Codex bot P1):Codex 指出当
tie_word_embeddings=True 且 get_input_embeddings() 返回包装器时,CausalMixin 会将包装器传给 tie_weights,但包装器没有 .weight 属性,导致 AttributeError。hmellor 回复 "Good catch, fixed",并随后在 causal.py 中改为遍历模块查找 VocabParallelEmbedding。
- 作用域确认(Isotr0py):Isotr0py 提醒
replace_embedding_class 只应用于文本主干嵌入,ViT 的视觉位置嵌入仍是 nn.Embedding(引用 clip.py)。hmellor 确认该方法只会在 model.get_input_embeddings 返回值上调用,因此不受影响。
- 组合模式设计意图(hmellor):hmellor 解释该 PR 支持的模式是
MyEmbedding 包装一个 nn.Embedding 并附加额外操作,模型通过 get_input_embeddings 返回该包装器;同时他决定收紧为“只允许单个 nn.Embedding”。
- 组合嵌入下 lm_head 权重绑定失败 (correctness): hmellor 承认并修复,改为遍历 modules() 查找 VocabParallelEmbedding 实例后再绑定。
- replace_embedding_class 的作用范围与视觉位置嵌入 (question): 确认作用范围仅限于文本主干嵌入,无需额外处理。
- 组合嵌入模式的设计意图与单一 nn.Embedding 限制 (design): 实现中通过 ValueError 强制单一 nn.Embedding,并补充了 test_replace_ambiguous_embedding 测试。
风险与影响
- 风险:
- 动态修改
__class__(embedding.__class__ = _rebase_on_vocab_parallel(...))属于高级 Python 技巧,若目标类存在 __slots__ 或 C 扩展类型可能失败;目前 nn.Embedding 子类普遍安全,但涉及第三方自定义嵌入时存在不可预见的兼容性风险。
_rebase_on_vocab_parallel 使用 lru_cache 缓存动态类,若缓存类被 GC 或涉及动态生成的类名冲突,可能产生行为异常;缓存键是原始类对象,风险较低。
VocabParallelEmbedding.__init__ 的 is_embedding_layer 判断改为 not isinstance(self, ParallelLMHead) 会影响所有嵌入和 lm_head 构造路径,但逻辑上更合理,回归风险有限。
can_replace_layer 从 type is 改为 isinstance 后,LoRA 可能尝试替换原本不应替换的派生类;虽然同时排除了 ParallelLMHead,但其他 VocabParallelEmbedding 子类(如未来新增)可能被意外匹配。
- 测试覆盖了内存中构造的场景,但未覆盖真实模型端到端(如 Qwen2VL 等组合嵌入模型)的 TP/PP 场景,存在集成测试盲区。
- 影响:影响范围:Transformers 建模后端(实验性后端)所有模型初始化路径;涉及 base.py、causal.py、LoRA 模块,以及 VocabParallelEmbedding 的量化方法判定逻辑。对用户影响:解决组合/包装输入嵌入模型(尤其是带缩放或额外变换的嵌入)无法在 Transformers 后端运行的问题,扩展了后端可支持模型范围。对团队影响:消除特判代码,为后续模型接入提供更稳健的基础设施;但动态 MRO 技巧需要维护者充分理解。
- 风险标记:动态类变更, 核心路径变更, 缺少端到端集成测试, LoRA 匹配放宽
关联脉络
- PR #49932 [Linear] [Kernel] add block-wise scaled_mm: 同为 vllm/model_executor 层改动,关注模型量化与嵌入层交互,且本 PR 在 VocabParallelEmbedding 中加入量化方法判定调整。
- PR #51249 [Bugfix][Model] Add missing fused_qkv_a_proj to Kimi-Linear packed_modules_mapping: 同样涉及模型嵌入与量化层映射的类替换逻辑,属于模型层正确性维护主题。
参与讨论