首页
/ Gemma 4 Assistant 使用指南:用 MTP 投机解码与 KV Sharing 加速 Gemma 4 推理

Gemma 4 Assistant 使用指南:用 MTP 投机解码与 KV Sharing 加速 Gemma 4 推理

2026-09-07 13:39:06作者:郦嵘贵Just

Gemma 4 Assistant 是一系列小型纯文本辅助模型(assistant model),配套 Gemma 4 主模型用于 Multi-Token Prediction(MTP)投机解码:由轻量辅助模型一次性草拟多个 token,再由主模型批量验证与接受,从而在保持输出质量的前提下降低整体解码时延。该模型架构于 2026-05-05 由社区贡献并合入本仓库(官方模型文档)。本文将以该文档为核心,结合本仓库中 configuration_gemma4_assistant.pymodeling_gemma4_assistant.py 的源码实现、以及生成管线与测试用例,完整讲解其架构原理、配置项与实战用法。

一、Gemma 4 Assistant 是什么

Gemma 4 Assistant 是一种 小型、纯文本 的辅助解码模型,专为 Gemma 4 主模型加速而设计。它本身不独立面向用户提供完整回答能力,而是作为 model.generate(..., assistant_model=...) 中的 草稿生成器(draft model / assistant model) 参与推理,实现 MTP 式的投机解码(speculative decoding)。

官方为 Gemma 4 各指令微调(IT)变体提供了预训练好的 Assistant 权重,覆盖 E2B、E4B、31B 与 26B-A4B(MoE) 四档参数规模,权重与主模型一同在 Gemma 4 官方发布中提供。

Gemma 4 Assistant 的训练/使用对象是 Gemma 4 主模型的文本解码阶段(其输入既支持纯文本,也支持多模态输入预处理后的文本序列),因此在 transformers 的 Pipeline 与 AutoModel 两条使用路径中,它都作为 assistant_model 参数出现。

二、架构核心:与 Gemma4TextModel 同源,但四点关键不同

Architecturally, Gemma 4 Assistant 复用了 Gemma 4 系列通用的 [Gemma4TextModel] 主干(backbone)。从源码看,Gemma4AssistantForCausalLM__init__ 中通过 AutoModel.from_config(text_config) 构造其文本主干,并叠加自身特有的 pre_projectionpost_projectionlm_head 与可选的 masked_embedding。与普通 Gemma 4 模型相比,它有以下四点关键差异:

1. 全模型使用 KV Sharing(KV 共享)

整个模型采用 KV sharing 技术。该技术最初随 Gemma 3n 引入,思想是让辅助模型直接复用目标主模型已经算好、缓存好的 KV cache

  • 由于 KV 内容由主模型填充,Assistant 可以 完全跳过 pre-fill(预填充)阶段——不需要对 prompt 再做一次注意力计算;
  • 前向过程中注意力计算量被大幅削减,因为目标 token 的 key/value 无需重复计算。

对应的实现约束在 Gemma4AssistantConfig.__post_init__ 中被强制落地:

# Assistant reuses the shared kvs across all layers to skip their calculation
# I.e. it acts as cache shared across the layers
if self.text_config is not None and not self.text_config.num_kv_shared_layers:
    self.text_config.num_kv_shared_layers = self.text_config.num_hidden_layers

即:只要用户没有显式指定,text_config.num_kv_shared_layers 会被自动置为 num_hidden_layers,表示每一层都共享主模型的 KV。配置校验 validate_architecture 还会进一步拒绝 num_kv_shared_layers != num_hidden_layers 的配置(报错信息为 “All layers in a Gemma 4 Assistant models must be shared.”),因此 "全层共享" 不是可选项,而是该模型架构的硬性前提。

2. position_ids 恒定

因为 KV cache 与主模型共享,而 Assistant 自身没有能力更新这段 cache,所以它只能在“固定位置”上进行预测:所有被草拟的 token 都使用同一个 position ID 参与注意力。这也是该模型被称为 "Single Position" 草稿生成器的原因(见下文生成管线部分)。

3. 输入 = 嵌入与隐状态的拼接,再经线性投影

为了适配静态 KV cache 与恒定 position_ids,Assistant 的输入并非常规的 input_ids 查表,而是取主模型 "最后看见的那个 token"embeddinghidden_states,将二者在最后一维拼接后,用一个 nn.Linear 投影到 Assistant 的模型空间。这一点在 Gemma4AssistantForCausalLM 中被编码为两个投影层:

self.pre_projection = nn.Linear(2 * self.backbone_hidden_size, self.hidden_size, bias=False)
self.post_projection = nn.Linear(self.hidden_size, self.backbone_hidden_size, bias=False)
  • pre_projection:把 [embedding; hidden_states](维度 2 × backbone_hidden_size)压回 assistant 自身主干使用的 hidden_size
  • post_projection:把 Assistant 主干输出的隐状态投影回主模型的 backbone_hidden_size 空间,供主模型作为下一轮验证输入。

值得注意的是,“最后看见的 token”在 assist decode 循环的不同阶段定义不同:

阶段 “最后看见的 token” 的定义
pre-fill 之后草拟的第一个 token prompt 的最后一个 token
同一草稿轮内随后的草拟步骤 本轮内 Assistant 自己刚生成的 token
两轮草稿之间 主模型上一轮验证并接受(accepted)的 token

该语义在生成器侧由主模型与 candidate generator 配合维护,Assistant 本身只负责“消费”传入的 inputs_embedsshared_kv_states

4. 使用 Cross-Attention 充分利用主模型上下文

为了最大化利用主模型的上下文,Assistant 在其注意力中引入 cross-attention:Assistant 自己生成的 query 状态可以直接去 attend 主模型共享 KV cache 中的 value 状态。这使得它在每轮草稿中能准确预测更多的 token(即提高单轮草稿长度与接受率),这是加速效果的核心来源。

forward 签名可以看到模型同时接收两类注意力来源(modeling 源码):

attention_mask: dict[str, torch.Tensor] | None = None,   # 键为 "full_attention" / "sliding_attention"
shared_kv_states: dict[str, tuple[torch.Tensor, torch.Tensor]] | None = None,

其中 shared_kv_states 是一个字典,包含 full_attention(全局注意力层)与 sliding_attention(滑动窗口注意力层)各自共享 KV 的 key/value,结构为 (key_states, value_states)forward 第一步就是强校验:

if inputs_embeds is None or shared_kv_states is None:
    raise ValueError("inputs_embeds and shared_kv_states cannot be None.")

这也说明:该模型无法脱离主模型独立前向,它必须以主模型产出的嵌入、隐状态与共享 KV 作为输入。

create_attention_masks 的掩码翻转技巧

由于 KV 是共享的、而 Assistant 的 position 恒定为 1(q_len == 1),常规因果掩码会出现方向歧义:滑动窗口注意力(SWA)在 q_idx=1kv>=1 时会被解释为“向前看未来 token”。为此,create_attention_masks 做了两件事:

  1. 在构造掩码前先沿 KV 轴翻转基础 attention mask,以保持位置无关的 padding 语义;
  2. 对滑动窗口掩码做“从未来视角翻转为过去视角”的处理(swa_mask = swa_mask.flip(dims=(-1,))),最终生成双向(bidirectional)掩码字典交给主干模型。

从源码注释看,当 q_len == 1 时全注意力与双向掩码没有区别(等价于全量关注),而滑动窗口部分则需要上述翻转修正。Gemma 4 主模型内部同样实现了 shared_kv_states 字典的产出路径(与 Gemma 3n 共享同一套机制),两者在接口上精确对应。

设备一致性处理

主模型与 Assistant 在多 GPU 场景下可以被拆分到不同设备。forward 中会主动做设备同步(源码 L170-L175):先把 inputs_embedsshared_kv_states 搬到 self.pre_projection.weight.device,计算完成后再把 last_hidden_statelogits 搬回 source_device。因此上层无需关心二者是否同卡。与之配套,模型类声明了 _skip_keys_device_placement = ["shared_kv_states"],避免框架重复搬移共享 KV。

三、与生成管线的集成:SinglePositionMultiTokenCandidateGenerator

Gemma 4 Assistant 的价值只有在 generate() 中才能体现。在本仓库的生成管线里,candidate generator 的选择逻辑如下:

# SinglePositionMultiTokenCandidateGenerator requires a target model that can provide, and an assistant model that
# can work from a shared_kv_states dictionary. Currently, the only models that can provide this are Gemma 3n and
# Gemma 4, and the only model that can work from it is a Gemma 4 Assistant
elif assistant_model is not None and assistant_model.__class__.__name__.startswith(
    ("Gemma4Assistant", "Gemma4UnifiedAssistant")
):
    if not self.__class__.__name__.startswith(("Gemma4", "Gemma3n")):
        raise ValueError(
            f"Expected class name to start with Gemma4 or Gemma3n. Got {self.__class__.__name__}."
            " Gemma4Assistant models require a target model that provides a shared_kv_states dictionary."
            " Currently, only Gemma4 and Gemma3n provide a shared_kv_states dictionary."
        )
    candidate_generator = SinglePositionMultiTokenCandidateGenerator(
        input_ids=input_ids,
        assistant_model=assistant_model,
        target_model_input_embeddings=self.get_input_embeddings(),
        ...
    )

可以提炼出三条关键约束,也是排错时最值得检查的点:

  1. assistant_model 的类名必须以 Gemma4AssistantGemma4UnifiedAssistant 开头,否则不会走该分支;
  2. 目标(主)模型类名必须以 Gemma4Gemma3n 开头,否则直接抛出 ValueError——因为只有这两类模型能对外提供 shared_kv_states 字典;
  3. 该分支使用 SinglePositionMultiTokenCandidateGenerator(定义于 candidate_generator.py),它对应"恒定 position + 每轮多 token 草稿"的生成模式,与 use_mtp=TrueMTPCandidateGenerator(主模型自带 MTP 头的情形)是两条不同的加速路径。

源码注释还可看到,此类 candidate generator “不会主动截断其草稿长度(如 MTP 总是草拟 num_mtp_layers 个 token)”,意味着草稿 token 数量受模型/配置决定而非启发式调度。因此官方文档将其描述为 “using the Multi-Token Prediction (MTP) method and associated candidate generator” 是准确的:一次前向草拟多个 token,再由主模型逐 token 验证接受

四、实战:两种调用方式

官方文档给出两种等效用法,都通过传入 assistant_model 把草稿模型接入生成。需要说明:官方文档注释提到 "generate text based on an image",其实现本质仍是文本解码阶段加速——示例中的图片是 Gemma 4 主模型的多模态输入,而文本生成由主模型与文本型 Assistant 协同完成。

方式一:Pipeline 一键使用

import torch
from transformers import pipeline

pipeline = pipeline(
    task="image-text-to-text",
    model="google/gemma-4-E2B-it",
    assistant_model="google/gemma-4-E2B-it-assistant",
)
pipeline(
    images="https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/pipeline-cat-chonk.jpeg",
    text="<|image|>\n\nWhat is shown in this image?"
)

pipeline(..., assistant_model=...) 会在内部自动完成主模型与 Assistant 的装配。这是最快的上手方式。

方式二:AutoModel 手动装配(推荐用于精细控制)

import torch
from transformers import AutoProcessor, AutoModelForImageTextToImage, AutoModelForCausalLM

model = AutoModelForImageTextToText.from_pretrained(
    "google/gemma-4-E2B-it",
    dtype=torch.bfloat16,
    device_map="auto",
)
assistant_model = AutoModelForCausalLM.from_pretrained(
    "google/gemma-4-E2B-it-assistant",
    dtype=torch.bfloat16,
    device_map="auto",
)

processor = AutoProcessor.from_pretrained(
    "google/gemma-4-E2B-it",
    padding_side="left"
)
messages = [
    {
        "role": "user", "content": [
            {"type": "image", "url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/pipeline-cat-chonk.jpeg"},
            {"type": "text", "text": "What is shown in this image?"},
        ]
    },
]
inputs = processor.apply_chat_template(
    messages,
    tokenize=True,
    return_dict=True,
    return_tensors="pt",
    add_generation_prompt=True,
).to(model.device)
input_len = inputs["input_ids"].shape[-1]

output = model.generate(**inputs, max_new_tokens=50, assistant_model=assistant_model)
print(processor.decode(output[0][input_len:], skip_special_tokens=True))

关键点解读:

  • 模型加载口径:主模型用 AutoModelForImageTextToText(多模态)或 AutoModelForCausalLM(纯文本);Assistant 统一用 AutoModelForCausalLM 加载,因为它本身就是纯文本 causal LM 风格模型,类名会解析为 Gemma4AssistantForCausalLM,从而满足上文生成管线中的前缀匹配条件。
  • padding_side="left":Gemma 4 / Assistant 体系支持批量左填充,与投机解码中共享 KV 的语义配合使用。
  • 聊天模板apply_chat_template(..., add_generation_prompt=True) 负责拼装 system/user 角色与图像占位 <|image|>,返回可直接送入 generate 的张量字典。
  • 解码max_new_tokens=50 限制新增 token 数;input_len 用于截取新增部分后再 decode
  • 若主模型未提供 shared_kv_states(即主模型类名不以 Gemma4/Gemma3n 开头),generate 会如前述那样直接报错提示。

纯文本场景的最小示例

forward 文档字符串中还给出了纯文本最小示例(modeling_gemma4_assistant.py):

from transformers import AutoTokenizer, Gemma4AssistantForCausalLM, Gemma4ForCausalLM

model = Gemma4ForCausalLM.from_pretrained("google/gemma-4-e2b-it")
assistant_model = Gemma4AssistantForCausalLM.from_pretrained("google/gemma-4-e2b-it-assistant")
tokenizer = AutoTokenizer.from_pretrained("google/gemma-4-e2b-it")

prompt = "What is your favorite condiment?"
inputs = tokenizer(prompt, return_tensors="pt")

# Generate
generate_ids = model.generate(inputs.input_ids, assistant_model=assistant_model, max_length=30)
tokenizer.batch_decode(generate_ids, skip_special_tokens=True)[0]

这一路径同样适用于 Gemma4TextModel 主模型(如 IT 文本变体)。

五、配置类:Gemma4AssistantConfig 参数详解

Gemma4AssistantConfig(定义在 configuration_gemma4_assistant.pymodel_type = "gemma4_assistant")通过 sub_configs 内嵌一个 text_config 子配置,指向其主干使用的 Gemma4TextConfig(默认按 gemma4_text 解析)。其自身暴露的顶层参数如下:

参数 默认值 含义
text_config None 主干模型配置,接受 Gemma4TextConfig 对象或 dict;若为 dict,按其中的 model_type 解析成正式配置对象
backbone_hidden_size 1536 Assistant 所配套的目标(主)模型的 hidden size;决定了 pre_projection 的输入宽度 2 × backbone_hidden_sizepost_projection 的输出宽度
use_ordered_embeddings False 是否使用为 Assistant 性能优化而重排的 embedding 表;若为 True,推理前需要把 embedding 顺序重排以与主模型对齐
num_centroids 2048 (有序 embedding 路径使用)质心(centroid)总数
centroid_intermediate_top_k 32 (有序 embedding 路径使用)激活的质心数量
tie_word_embeddings True 是否将 lm_head.weight 与主干 embedding 权重绑定

关于 use_ordered_embeddings

当该选项开启时,模型不再使用完整稠密 lm_head,而是采用 Gemma4AssistantMaskedEmbedder实现源码)做 centroid + top-k 掩码词表投影

  1. 先用一个无 bias 的 nn.Linear(hidden_size, num_centroids) 计算每个 token 落在 2048 个质心上的 logits;
  2. 取 top-32 质心,利用与 lm_head 权重对齐的 token_ordering 缓冲区拿到每个质心对应的 canonical token 位置;
  3. 仅在这些位置上与 lm_head 权重做点积,其余位置填 mask_value(最小 logit − 1),最后通过 scatter_ 还原为完整词表大小的 logits。

这样可把每步 LM head 计算限制在 top_k × (vocab_size / num_centroids) 个 token 上,是大词表下进一步省算力的手段;默认关闭时则退化为常规 self.lm_head(last_hidden_state)。官方将 checkpoint 的该类配置在文档开头标注为 "google/gemma-4-e2b-it"(默认 checkpoint 示例),加载实际权重时以对应 checkpoint 的 config.json 为准。

配置约束(validate_architecture)

为了与共享 KV / 固定位置 / 拼接输入等架构假设一致,validate_architecture 会对 text_config 强校验以下条件,不满足即抛错:

  • hidden_size_per_layer_input 必须为 0(不能用 Gemma 4 的 Per-Layer Embeddings 逐层输入);
  • enable_moe_block 必须为 False(Assistant 主干不启用 MoE block);
  • use_double_wide_mlp 必须为 False
  • vocab_size_per_layer_input 必须为 0
  • num_kv_shared_layers 必须等于 num_hidden_layers(全层 KV 共享)。

手动构造配置示例

from transformers import Gemma4AssistantConfig, Gemma4TextConfig

# 构造一个 Gemma 4 Text 配置(此处仅示意,实际建议直接从 checkpoint 加载)
text_config = Gemma4TextConfig(hidden_size=..., num_hidden_layers=..., vocab_size=...)

# 用其构造 Gemma 4 Assistant 配置
configuration = Gemma4AssistantConfig(text_config)

# 从配置初始化模型
model = Gemma4AssistantForCausalLM(configuration)

注意实际部署时更推荐直接 from_pretrained("google/gemma-4-E2B-it-assistant"),让配置随 checkpoint 一起加载,避免手写 text_config 时出现与权重不一致的问题。

六、模型类与 I/O 形态:Gemma4AssistantForCausalLM

Gemma4AssistantForCausalLM源码)继承 Gemma4AssistantPreTrainedModel 并混入 GenerationMixin,其 forward 输入输出形态与常规 CausalLM 差异明显:

输入参数 说明
input_ids 不使用,仅为兼容签名而保留,调用时会被忽略
inputs_embeds 必须提供:主模型最后一个 token 的 embedding 与 hidden_states 的拼接(由 candidate generator 准备)
position_ids 恒定值(单位置预测)
attention_mask 字典类型:{"full_attention": ..., "sliding_attention": ...} 的双向掩码
shared_kv_states 必须提供:主模型各层共享 KV,dict[str, tuple(key, value)]
use_cache 不使用,保留仅为兼容签名

输出则为 Gemma4AssistantOutput(继承 BaseModelOutput),包含:

  • logits:形状 (batch_size, sequence_length, config.vocab_size) 的 LM head 预测分数(SoftMax 前);
  • last_hidden_state:已经过 post_projection 投影回 backbone_hidden_size 的隐状态,供主模型继续使用;
  • 其余 hidden_statesattentions 透传主干模型输出。

模型还声明了与分布式/编译能力相关的标志:_supports_flash_attn = True_supports_sdpa = True_can_compile_fullgraph = True_supports_attention_backend = True 以及 supports_gradient_checkpointing = True,并对 lm_head 配置了 FSDP(keep_full_weight)、张量并行(colwise_gather_output)与流水线并行(["hidden_states"] -> ["logits"])策略,便于在大规模部署下按需启用。

七、测试与质量保障

本仓库为 Gemma 4 Assistant 提供了集成测试(test_modeling_gemma4_assistant.py),测试用例如下:

  • 模型名使用 google/gemma-4-E2B-it,Assistant 使用 google/gemma-4-E2B-it-assistant
  • 图文场景 test_model_with_image:加载 Gemma4ForConditionalGeneration 与 Assistant,经 apply_chat_template 处理 system/user 消息与图片,model.generate(..., assistant_model=assistant, max_new_tokens=30, do_sample=False) 输出对画面内容的描述;
  • 纯文本场景 test_model_text_only:加载 AutoModelForCausalLM 与 Assistant,让用户指令(如 "Write a poem about Machine Learning.")走相同装配路径,校验 do_sample=False 下的确定性输出。

这些测试标注为 @slow 且当前版本中带 @unittest.skip(reason="Update after release") 跳过标记(等待官方权重随发布更新后启用),但代码本身验证了两个要点:Assistant 必须与 Gemma4/Gemma3n 系列主模型搭配,且 主模型侧 generate 会负责解码与校验,Assistant 只贡献草稿。若你想自行复现,需要 GPU 与可访问官方 checkpoint 的网络环境。

八、适用场景与前提小结

把官方文档与本仓库源码对照,使用 Gemma 4 Assistant 前请确认以下几点,可避免绝大多数踩坑:

  1. 主模型必须是 Gemma 4(含文本/多模态变体)或 Gemma 3n:只有它们的前向会产出 shared_kv_states 字典,Assistant 才能跳过 pre-fill 并共享 KV;
  2. 主模型与 Assistant 的权重配套:Assistant checkpoint 与主模型按同一代(IT 变体、同参数量级)发布,混用不同规模的 checkpoint 不具备事实依据且必然出错;
  3. 输入规范inputs_embedsshared_kv_states 缺一不可,模型不读取 input_ids
  4. 面向推理加速:该模型设计目标是解码吞吐/时延优化(通过单轮多 token 草稿 + 主模型验证),具体收益与批大小、生成长度、草稿接受率强相关,请基于自身硬件实测,不宜一概而论;
  5. 参考主模型文档:更完整的 Gemma 4 能力(多模态、函数调用、上下文窗口、处理器用法)参见 Gemma4 模型文档;KV sharing 的底层机制可追溯 Gemma 3n 模型文档;仓库中另有面向多模态统一主模型的 Gemma 4 Unified Assistant 实现 与文档 gemma4_unified_assistant.md,采用同样的加速思路。

总而言之:Gemma 4 Assistant 是 Transformers 投机解码生态中“以共享 KV + 单位置多 token 草稿”换取解码效率的代表性实现。理解其四点架构差异与生成管线的装配约束后,你就能把它正确地接入自己的 Gemma 4 推理服务,并在源码层面定位加速链路中的各类问题。

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