Gemma 4 Assistant 使用指南:用 MTP 投机解码与 KV Sharing 加速 Gemma 4 推理
Gemma 4 Assistant 是一系列小型纯文本辅助模型(assistant model),配套 Gemma 4 主模型用于 Multi-Token Prediction(MTP)投机解码:由轻量辅助模型一次性草拟多个 token,再由主模型批量验证与接受,从而在保持输出质量的前提下降低整体解码时延。该模型架构于 2026-05-05 由社区贡献并合入本仓库(官方模型文档)。本文将以该文档为核心,结合本仓库中 configuration_gemma4_assistant.py 与 modeling_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_projection、post_projection、lm_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" 的 embedding 与 hidden_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_embeds 与 shared_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=1 且 kv>=1 时会被解释为“向前看未来 token”。为此,create_attention_masks 做了两件事:
- 在构造掩码前先沿 KV 轴翻转基础 attention mask,以保持位置无关的 padding 语义;
- 对滑动窗口掩码做“从未来视角翻转为过去视角”的处理(
swa_mask = swa_mask.flip(dims=(-1,))),最终生成双向(bidirectional)掩码字典交给主干模型。
从源码注释看,当 q_len == 1 时全注意力与双向掩码没有区别(等价于全量关注),而滑动窗口部分则需要上述翻转修正。Gemma 4 主模型内部同样实现了 shared_kv_states 字典的产出路径(与 Gemma 3n 共享同一套机制),两者在接口上精确对应。
设备一致性处理
主模型与 Assistant 在多 GPU 场景下可以被拆分到不同设备。forward 中会主动做设备同步(源码 L170-L175):先把 inputs_embeds 与 shared_kv_states 搬到 self.pre_projection.weight.device,计算完成后再把 last_hidden_state 与 logits 搬回 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(),
...
)
可以提炼出三条关键约束,也是排错时最值得检查的点:
assistant_model的类名必须以Gemma4Assistant或Gemma4UnifiedAssistant开头,否则不会走该分支;- 目标(主)模型类名必须以
Gemma4或Gemma3n开头,否则直接抛出ValueError——因为只有这两类模型能对外提供shared_kv_states字典; - 该分支使用
SinglePositionMultiTokenCandidateGenerator(定义于 candidate_generator.py),它对应"恒定 position + 每轮多 token 草稿"的生成模式,与use_mtp=True走MTPCandidateGenerator(主模型自带 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.py,model_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_size 与 post_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 掩码词表投影:
- 先用一个无 bias 的
nn.Linear(hidden_size, num_centroids)计算每个 token 落在 2048 个质心上的 logits; - 取 top-32 质心,利用与
lm_head权重对齐的token_ordering缓冲区拿到每个质心对应的 canonical token 位置; - 仅在这些位置上与
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_states、attentions透传主干模型输出。
模型还声明了与分布式/编译能力相关的标志:_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 前请确认以下几点,可避免绝大多数踩坑:
- 主模型必须是 Gemma 4(含文本/多模态变体)或 Gemma 3n:只有它们的前向会产出
shared_kv_states字典,Assistant 才能跳过 pre-fill 并共享 KV; - 主模型与 Assistant 的权重配套:Assistant checkpoint 与主模型按同一代(IT 变体、同参数量级)发布,混用不同规模的 checkpoint 不具备事实依据且必然出错;
- 输入规范:
inputs_embeds、shared_kv_states缺一不可,模型不读取input_ids; - 面向推理加速:该模型设计目标是解码吞吐/时延优化(通过单轮多 token 草稿 + 主模型验证),具体收益与批大小、生成长度、草稿接受率强相关,请基于自身硬件实测,不宜一概而论;
- 参考主模型文档:更完整的 Gemma 4 能力(多模态、函数调用、上下文窗口、处理器用法)参见 Gemma4 模型文档;KV sharing 的底层机制可追溯 Gemma 3n 模型文档;仓库中另有面向多模态统一主模型的 Gemma 4 Unified Assistant 实现 与文档 gemma4_unified_assistant.md,采用同样的加速思路。
总而言之:Gemma 4 Assistant 是 Transformers 投机解码生态中“以共享 KV + 单位置多 token 草稿”换取解码效率的代表性实现。理解其四点架构差异与生成管线的装配约束后,你就能把它正确地接入自己的 Gemma 4 推理服务,并在源码层面定位加速链路中的各类问题。
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 StartedRust0626
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00