手搓 EmbeddingGemma

一位开发者近日用 PyTorch 从零复现 EmbeddingGemma,拆开了分词、RoPE、GQA、RMSNorm、注意力池化与 Matryoshka 向量截断等关键组件。本文结合模型结构,解释它为什么适合本地 RAG,以及复现时最容易踩中的坑。
手搓 EmbeddingGemma:一个嵌入模型到底是怎么工作的
一位开发者近日发布了《Implementing Embedding Gemma from scratch in PyTorch》,尝试不依赖现成 Transformer 封装,而是用原生 PyTorch 拆解并复现 EmbeddingGemma 的核心推理路径。这篇文章的价值不在于又造了一个向量模型,而在于把一个看似只有“文本进、向量出”的黑盒,拆成了分词、位置编码、归一化、分组查询注意力、池化和向量压缩等一组可以逐个验证的部件。
EmbeddingGemma 是 Google 基于 Gemma 3 架构训练的多语言文本嵌入模型,参数量约为 3.08 亿,主要用于语义检索、相似度搜索、分类、聚类和本地 RAG。它的定位很明确:不是和大语言模型比生成能力,而是用更小的体积,把“这两段文本是不是在说同一件事”计算得更快、更便宜。

先说结论:复现难点不在 Transformer 主干
从零实现 EmbeddingGemma,最容易产生的误解是“把一个 Decoder-only Transformer 写出来就完成了”。实际上,Transformer 主干只是第一关,最终向量是否能与官方模型对齐,还取决于输入模板、tokenizer 行为、位置编码、padding 处理、池化策略和归一化顺序。
换句话说,复现一个嵌入模型更像复刻一条生产线,而不是重新组装一台发动机。发动机能转,并不代表产品能用;某个特殊 token 少了一个、attention mask 偏移一位,最后的余弦相似度就可能完全失真。
Embedding 模型是把文本映射成固定长度数值向量的模型,向量之间的距离用于表示文本在语义空间中的接近程度。对开发者而言,它最常见的用途不是直接回答问题,而是先把文档和查询变成向量,再从向量库中找出可能相关的内容。
EmbeddingGemma 的主要特点可以概括为:
- 约 3.08 亿参数,基于 Gemma 3 架构。
- 支持 100 多种语言,面向跨语言检索和语义匹配场景。
- 支持通过 Matryoshka Representation Learning 调整输出维度。
- 量化后可以在低于 200MB RAM 的设备上运行。
- 适合手机、笔记本电脑、平板电脑和离线应用。
- 输出向量主要用于检索、相似度、分类和聚类,而不是文本生成。
它的实际意义是:过去需要服务器调用的文档搜索,现在可以放到本地设备完成。用户的笔记、邮件或企业内部资料不必离开设备,断网时也能继续完成基础检索。
第一层:输入格式比想象中更重要
EmbeddingGemma 的输入不是简单地把原始字符串丢进模型。嵌入模型通常会区分查询和文档,因为“我要找什么”和“被搜索的内容”在训练时承担的角色不同。
例如,一个检索任务可以把输入抽象成两类:
查询:task: search query | query: 如何在手机上运行本地 RAG
文档:task: search result | title: 移动端 RAG 部署指南 | text: ...
这里的前缀不是装饰性文本,而是模型训练目标的一部分。它相当于告诉模型:“现在这段输入是用户的问题”或者“现在这段输入是候选文档”。如果把文档和查询都用同一种模板编码,模型仍然可能返回向量,但检索效果会出现明显下降。
从零复现时,第一步应当确认官方模型卡中规定的任务前缀、字段顺序、分隔符和最大长度。不能只看模型结构文件,因为输入模板经常位于模型卡、推理示例或配套代码中,而不是 Transformer 配置里。
Tokenizer 也不能被当作普通字符串切分器。它负责把文本映射为 token ID,同时处理特殊 token、空格、Unicode 字符和 padding。中文、英文、数字、代码混排时,token 数量并不等于字符数量;一段看起来只有几百字的技术文档,可能因为标点、路径和英文标识符变成更多 token。
一个实用的检查方法是,先用官方 tokenizer 对固定样本编码,记录以下结果:
- 输入 token ID 的完整序列。
- BOS、EOS 和 padding token 的位置。
- attention mask 的形状和值。
- 截断前后的 token 数量。
- 查询模板和文档模板之间的差异。
只要这些结果没有对齐,后面即使把注意力层写得完全正确,也很难得到一致的向量。
第二层:Gemma 风格 Transformer 的核心组件
Transformer 是一种通过注意力机制处理序列的神经网络架构,它不需要像传统循环网络那样逐字处理文本,而是可以同时比较序列中不同位置的 token。
EmbeddingGemma 的主干继承了 Gemma 3 系列的设计思路。开发者在 PyTorch 中重点拆解了以下组件:RMSNorm、旋转位置编码 RoPE、分组查询注意力 GQA、前馈网络和残差连接。
RMSNorm:只做尺度校准
RMSNorm 是一种归一化方法,它使用向量均方根调整激活值的尺度,不计算均值,因此比 LayerNorm 少了一步中心化操作。
可以把它理解为给每一层的信号装上自动增益控制器:信号太大时压低,太小时放大,但不主动改变信号的方向。一个简化实现如下:
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
rms = torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
return x * rms * self.weight
这里的 eps、权重初始值以及归一化发生在注意力和前馈网络之前还是之后,都必须与目标实现一致。工程上最常见的错误不是公式写错,而是把 Pre-Norm 和 Post-Norm 的顺序搞反。
RoPE:让注意力知道 token 的位置
RoPE 是旋转位置编码,它通过对查询向量和键向量做与位置有关的二维旋转,把相对位置信息注入注意力计算。
普通的 token embedding 只知道“这个词是什么”,RoPE 则补充“这个词在序列中的什么位置”。对检索模型来说,位置仍然重要:同一组词出现在标题、句首或正文末尾,语义贡献可能不同。
复现 RoPE 时需要特别检查三件事:
- 旋转频率的计算方式。
- head dimension 是否能被旋转维度正确拆分。
- cos、sin 张量是否与 query、key 的 dtype 和设备一致。
如果使用半精度推理,位置编码缓存最好按照实际设备和 dtype 管理,否则可能出现隐式类型转换,既影响速度,也可能导致和官方结果之间出现不易察觉的误差。
GQA:用更少的 KV 头降低内存访问
GQA,即 Grouped-Query Attention,是让多个 query head 共享较少 key/value head 的注意力结构。它保留了多头查询的表达能力,同时减少 KV 缓存和内存读写。
在生成模型中,GQA 的优势通常体现为更低的 KV cache 成本;在嵌入模型中,它更直接的收益是减少推理过程中的中间张量规模。对于手机和笔记本这类内存受限设备,这比单纯减少几层网络更有实际价值。
实现时不能把 query、key、value 的头数混为一谈。假设 query 有 8 个头、key/value 有 4 个头,那么每个 KV 头要被对应的 query 组复用 2 次。复用方式应通过 reshape 或 expand 明确表达,避免错误地复制整个序列张量。
注意力的基本计算仍然是:
Attention(Q, K, V) = softmax(QK^T / sqrt(d))V
但在实际模型中,还需要加入 RoPE、attention mask、头数映射以及可能的缩放约定。任何一个细节不一致,都会使最终向量偏离官方实现。
第三层:从 token 表示得到句向量
池化是把一串 token 表示压缩为一个句向量的过程。它决定模型最终保留哪些信息,也是复现 EmbeddingGemma 时最容易被低估的环节。
常见做法包括取最后一个 token、取特殊 token、对所有有效 token 做平均,或者使用经过训练的池化层。对于句子嵌入模型,平均池化通常更直观:把所有有效 token 的隐藏状态求平均,再通过投影和归一化得到固定长度向量。
但平均池化有一个硬性要求:必须排除 padding。错误示例是直接对整个序列求平均;当 batch 内文本长度不一致时,padding 的零值或对应 embedding 会改变句向量方向。
带 mask 的平均池化可以写成:
def mean_pool(hidden_states, attention_mask):
mask = attention_mask.unsqueeze(-1).to(hidden_states.dtype)
summed = (hidden_states * mask).sum(dim=1)
count = mask.sum(dim=1).clamp_min(1e-6)
return summed / count
之后通常还要做 L2 归一化。L2 归一化是把向量除以自身的欧氏范数,使所有向量落在单位球面上。这样计算点积时,结果就等价于余弦相似度,向量长度不会因为文本长短或激活尺度而主导排序。
def normalize_embeddings(x):
return x / x.norm(dim=-1, keepdim=True).clamp_min(1e-12)
需要注意,池化和归一化的顺序不能随意调换。正确顺序应以官方实现为准:先对 token 表示进行规定的池化或投影,再做最终归一化。只看输出维度而不检查归一化方式,很容易得到“形状正确、排序错误”的结果。
Matryoshka:一条向量,多个尺寸
Matryoshka Representation Learning,简称 MRL,是一种让同一个向量的前缀维度也具备语义能力的训练方法。它的核心思想是:完整向量可以用于高精度检索,但只取前 256 维、128 维等较短前缀,也仍然能保持可用的排序效果。
这对向量数据库尤其重要。向量维度越高,索引占用的内存和计算成本通常越高。假设数据库中有 100 万条向量:
| 向量维度 | FP32 原始存储 | FP16 原始存储 | 适合场景 | |---:|---:|---:|---| | 768 | 约 2.93GB | 约 1.46GB | 追求召回率的离线检索 | | 512 | 约 1.95GB | 约 0.98GB | 通用知识库 | | 256 | 约 0.98GB | 约 0.49GB | 移动端和资源受限设备 | | 128 | 约 0.49GB | 约 0.24GB | 粗筛、轻量相似度匹配 |
上述数字只计算向量本身,没有包含索引结构、元数据和数据库额外开销。实际节省会因索引类型不同而变化,但量级关系是明确的。
MRL 的使用方式不是重新训练一个 256 维模型,而是先生成完整向量,再截取前 N 维并重新归一化:
def truncate_and_normalize(x, dim):
x = x[..., :dim]
return x / x.norm(dim=-1, keepdim=True).clamp_min(1e-12)
“重新归一化”是关键。截断后向量的长度已经改变,如果直接拿未归一化的前缀计算点积,结果不能与完整向量的相似度直接比较。
不过,MRL 不是免费午餐。低维前缀通常会损失一部分召回率,尤其是长文档、细粒度技术问答和相近概念较多的语料。比较稳妥的做法是先在自己的数据集上测试 768、512、256 和 128 维,再决定索引配置,而不是看到维度更小就直接切换。
与其他嵌入模型怎么选
EmbeddingGemma 的优势是体积、语言覆盖和本地部署平衡得比较好,但它并不是所有场景下的最佳选择。模型选择应当同时看参数量、向量维度、上下文长度、语言覆盖、许可证和硬件成本。
| 模型 | 参数量或定位 | 向量维度 | 多语言能力 | 主要优势 | 更适合的场景 | |---|---:|---:|---|---|---| | EmbeddingGemma | 约 3.08 亿 | 支持 MRL,常用完整维度为 768 | 100+ 语言 | 小体积、可量化、可本地运行 | 本地 RAG、移动端、隐私敏感数据 | | BGE-M3 | 约 5.68 亿 | 1024 | 多语言 | 稠密、稀疏和多向量检索能力较完整 | 中文知识库、混合检索 | | multilingual-e5-large | 约 5.6 亿 | 1024 | 多语言 | 社区使用广、检索范式成熟 | 跨语言搜索和通用语义匹配 | | MiniLM 类模型 | 通常低于 1 亿 | 常见为 384 | 取决于具体版本 | 推理速度快、占用低 | 低延迟分类、边缘设备粗筛 |
表中的参数量和维度不能直接等价为效果。嵌入模型的训练数据、任务模板和评测集同样重要。尤其对中文技术资料,英文或多语言 MTEB 分数并不能完全预测中文产品搜索的实际效果。
EmbeddingGemma 的判断可以简单概括为:如果优先考虑离线运行、安装包大小和隐私,它很有竞争力;如果需要极强的中文领域效果,或者依赖稀疏检索与多向量检索,BGE-M3 这类模型仍然值得优先评估;如果只追求最低延迟,参数更小的 MiniLM 类模型可能更划算。
从零复现的验证顺序
开发者不应一上来就用整套 MTEB 评测。更高效的验证方式是分层对齐,每一层只检查一个问题。
1. 先对齐 tokenizer
选择 5 到 10 条固定文本,覆盖中文、英文、数字、URL、代码和空字符串。逐项比较 token ID、特殊 token 和 attention mask。只要这里不一致,就暂停后续排查。
2. 再对齐单层输出
将随机输入送入官方模型和自写模块,比较 embedding 层、RMSNorm、RoPE、注意力输出和前馈网络输出。使用 torch.testing.assert_close 时,先用宽松容差定位数量级问题,再逐步收紧容差。
3. 最后对齐句向量
对一组语义相近和语义无关的句子计算余弦相似度。不要只比较两个向量的绝对差异,还要比较排序结果,因为检索系统真正关心的是相关文档能否排在前面。
4. 测试 padding 和 batch
单条文本能跑通,不代表批量推理正确。将同一文本分别以单条和不同长度 batch 输入,检查最终向量是否一致。如果 batch 大小改变后向量明显变化,通常意味着 padding mask、池化或位置索引处理存在问题。
5. 检查量化影响
EmbeddingGemma 的低内存优势很大程度来自量化。量化是用更低位宽的数据类型近似保存模型权重,从而降低内存占用和带宽压力。INT8、4-bit 量化都可能影响相似度排序,必须在目标数据集上测 Recall@K、MRR 或 nDCG,而不是只看模型是否能成功加载。
实际部署时,别只盯着模型大小
EmbeddingGemma 量化后低于 200MB RAM 的宣传,对移动端部署很有吸引力,但“模型文件小”不等于“应用运行只占这么多内存”。Tokenizer、运行时、临时激活、batch 缓冲区和向量索引都要占用资源。
在本地 RAG 中,通常应把流程拆成两个阶段:文档入库阶段离线生成向量,查询阶段实时生成用户问题向量。文档向量可以使用较大的 batch;查询向量则更看重首 token 延迟和峰值内存。对手机应用来说,batch size 为 1 或 2 往往比盲目追求吞吐更合理。
分块策略也会直接影响模型表现。把一篇文档切成固定字符数并不一定好。更可靠的做法是优先按标题、段落和列表切分,并保留少量上下文重叠。过大的 chunk 会把多个主题混在一个向量中,过小的 chunk 则会丢失指代关系和条件约束。
Embedding 模型也不能替代重排模型。一个常见的高质量 RAG 流程是:先用 EmbeddingGemma 对所有候选文档做向量召回,再用更精确的交叉编码器或大模型重排前几十条结果。嵌入模型负责快速缩小搜索范围,重排模型负责判断查询和文档之间的细粒度关系。
这次复现最值得学的是什么
这篇 PyTorch 从零复现的真正价值,是让开发者看到嵌入模型的效果来自一连串精确但不神秘的工程选择:输入模板决定任务语义,Tokenizer 决定序列边界,RoPE 注入位置信息,GQA 控制注意力成本,池化把 token 表示变成句向量,L2 归一化让距离可比较,MRL 则在效果和存储之间提供可调旋钮。
它也提醒了一个容易被忽视的事实:模型结构相同,不代表模型行为相同。对于 EmbeddingGemma,最值得复现的不是“能不能跑出 768 个浮点数”,而是同一批文本在向量空间中的相对关系能否稳定复现。
截至 2026 年 9 月 5 日,EmbeddingGemma 更像是一个面向本地 AI 应用的基础设施组件,而不是单纯的模型发布。它适合被嵌入搜索框、笔记库、离线知识助手和端侧推荐系统。对开发者来说,最现实的路径是先复现单条推理,再对照官方权重验证数值,最后用自己的中文或垂直领域数据决定输出维度和量化方案。
如果只是想快速接入,直接使用官方权重和成熟推理框架更省时间;如果要做端侧优化、模型移植、量化适配或排查检索效果问题,那么这次从零实现提供了一条非常清晰的学习路径。EmbeddingGemma 的门槛不在模型规模,而在那些决定最终排序质量的细节。



