gemma_pytorch
Google 官方 PyTorch 实现 Gemma 大模型,支持 GPU/TPU/CPU 多硬件,
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
Google 官方 PyTorch 实现 Gemma 大模型,支持 GPU/TPU/CPU 多硬件,
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
2024 年 2 月,Google 宣布开源 Gemma——一个基于 Gemini 技术的小型大语言模型家族,首批开放了 2B 和 7B 两种参数规模。与当时风头正劲的 Meta Llama 2 不同,Gemma 采用了完全不同的实现路线:它没有发布 PyTorch 原生模型,而是以 JAX/Keras 格式发布。对于习惯了 PyTorch 生态的开发者来说,这意味着需要额外的转换步骤才能使用。
Google 显然也意识到了这一点。google/gemma_pytorch 仓库正是官方给出的"标准答案"——一套完整的 PyTorch 实现,覆盖 Gemma v1、v2、v3、CodeGemma 全部版本,让全球数百万 PyTorch 开发者可以直接上手这款 Google 出品的顶级开源模型。
图1:Google 官方开源项目 logo
要理解 Gemma 的定位,需要从它的"母体"说起。Gemini 是 Google 最强大的多模态大模型,训练时使用了 TPU v5 超级计算机、数万块芯片和海量数据。Gemma 的目标是将 Gemini 的核心技术——尤其是 Transformer 架构的改进版本——浓缩到一个普通研究者也能运行的规模。
Gemma 有几个显著特点让它区别于同期的开源模型:
架构创新:Gemma 采用了与 Llama 类似的 Transformer 基础架构,但引入了多项来自 Gemini 的改进,包括 Grouped-Query Attention(GQA,分组查询注意力)和 RoPE 位置编码的变体,在长上下文处理上表现更稳定。
训练数据透明:与很多开源模型的模糊训练数据描述不同,Google 明确表示 Gemma 预训练数据来自网页文档、代码、科学论文等,并通过 Common Crawl 清洗,数据量约为 6 万亿 tokens。
多版本覆盖:Gemma 3 支持多模态(图像+文本),27B_v3 版本是目前最强的开源 Gemma 模型,在多项 benchmark 上接近 Llama 3 70B 的水平。
许可友好:Gemma 采用 Apache 2.0 许可,这是业界最宽松的开源许可之一,允许商业使用和修改,无需支付版税。
google/gemma_pytorch 仓库的功能远不止"把 JAX 转成 PyTorch"这么简单。深入看代码,你会发现这是一套经过精心设计的推理框架:
仓库同时支持三种硬件后端:
这个设计体现了 Google 的战略意图——让 Gemma 不仅能在消费级 GPU 上跑,还能无缝对接 Google Cloud TPU 资源。这对于学术研究机构和企业来说意义重大:TPU 的性价比在大规模推理场景下往往更高。
仓库实现了 Gemma 全系列模型的 PyTorch 版本:
每种尺寸都提供预训练(base)和指令微调(instruct)两个版本。
Gemma 使用 Google 自研的 SentencePiece tokenizer(不同于 Llama 的 tiktoken),每个模型的 tokenizer 独立,词汇表大小约 256K tokens。仓库在 tokenizer/ 目录下提供了完整的 tokenizer 实现,支持多语言和代码 tokenization。
scripts/run.py 提供了开箱即用的推理脚本,支持:
scripts/run_multimodal.py 则专门处理 Gemma 3 多模态版本,可以同时输入图像和文本。
最简单的方式是 pip install,配合 HuggingFace 下载模型权重:
pip install sentencepiece transformers torch
huggingface-cli download google/gemma-3-4b-it-pytorch
之后用几行 Python 即可推理:
from gemma import config, model
import torch
model_path = "/path/to/your/model"
model_cfg = config GemmaConfig(variant="2b-it")
with model.CausalModel(model_cfg) as m:
m.load_weights(model_path)
output = m.generate(...)
仓库提供了三个官方 Dockerfile:
Dockerfile:基于 pytorch/pytorch:2.1.2-cuda11.8,适合 NVIDIA GPUxla.Dockerfile:基于 Google Cloud TPU 镜像,适合 TPU 训练xla_gpu.Dockerfile:PyTorch/XLA + CUDA 组合,适合大规模推理Docker 方式的优势是环境隔离,不依赖本地 CUDA 版本。
| 模型规模 | 最低 VRAM | 推荐 VRAM | 推理精度 |
|---|---|---|---|
| 1B | 2GB | 4GB | INT8/F16 |
| 4B | 8GB | 12GB | F16 |
| 7B | 14GB | 16GB+ | F16 |
| 27B | 56GB | 80GB | F16/INT4 |
27B 模型即使使用 INT4 量化,也需要至少 16-20GB VRAM,普通的 RTX 3090/4090(24GB)勉强可以跑,但建议用 A100 或 H100。
Gemma 的出现对开源大模型生态有深远影响。在此之前,开源 LLM 的主力是 Meta(Llama)和 Mistral。Google 的入局带来了几个不同:
可信度:Google 作为 Transformer 论文的原创者之一,在 LLM 领域的技术积累毋庸置疑。Gemma 的技术报告质量明显高于一般开源项目。
多模态优先:Gemma 3 的多模态版本在开源模型中属于较早一批,而 PyTorch 实现让开发者可以更灵活地定制视觉编码器(如替换 SigLIP 视觉塔)。
Google 生态绑定:TPU 支持、Colab 教程、Kaggle 模型发布,Google 正在构建一个从训练到推理的完整开源生态。
从趋势看,Gemma 正在走一条不同于 Llama 的路——不追求"最大",而是追求"最适合特定硬件平台的高效模型"。这对 AI 应用开发者是好事,意味着可以更低成本获得接近 GPT-4 水平的能力。
google/gemma_pytorch 是 Google 官方发布的 PyTorch 实现库,覆盖 Gemma 全系列模型(1B27B)和全版本(v1v3),支持 GPU/TPU/CPU 三种硬件后端。适合学术研究和企业部署,但需要一定 Python 能力,无开箱即用的 Web UI。随着 Gemma 3 多模态版本的加入,这套代码库正成为开源多模态 AI 的重要基础设施之一。