Gemma
Реализация:
llm/src/llm/models/gemma/gemma.py· классGemmaНоутбук:notebooks/gemma.ipynb
Место в линейке: развивает ту же базу (RoPE + RMSNorm), что и LLaMA/Mistral, но с собственным вариантом attention и FFN — не входит в основную цепочку GPT → Mixtral. Это последняя глава части II; дальше — глоссарий и оглавление.
Что вы узнаете
Заголовок раздела «Что вы узнаете»- Какие решения отличают Gemma от LLaMA: GeGLU, MQA в модели 2B, словарь 256k с общей матрицей эмбеддингов и выхода, масштаб эмбеддингов , RMSNorm с множителем .
- Почему у Gemma 7B и как это выражается ключом
head_size. - Как записать прямой проход Gemma в формулах и какие ключи конфига делают модель такой же, как в статье.
- Как посчитать параметры 2B и 7B и сверить их с таблицей статьи.
- Как загрузить веса HuggingFace и почему к весам RMSNorm прибавляется 1.
Предварительные знания
Заголовок раздела «Предварительные знания»- LLaMA — общая основа RoPE + RMSNorm + gated FFN.
- Механизм внимания — MQA как частный случай GQA.
- Эмбеддинги и выходная проекция — weight tying и масштаб √d.
- Feed-forward сеть и активации — GeGLU.
Gemma (Google DeepMind, 2024, arXiv:2403.08295) в этом репозитории реализована как RoPE + RMSNorm трансформер с Multi-Query Attention по умолчанию (MQA — одна общая голова K/V на все Q-головы, предельный случай GQA; число K/V-голов задаётся ключом num_kv_heads) и GeGLU-FFN (GELU-gated, а не SiLU-gated, как в SwiGLU). Ключи конфига из таблицы ниже делают модель такой же, как Gemma 2B/7B, вплоть до загрузки весов HuggingFace.
По сравнению с LLaMA у Gemma четыре заметных отличия: GeGLU вместо SwiGLU, MQA в модели 2B, словарь 256k с общей матрицей эмбеддингов и выходной проекции, умножение эмбеддингов на . Ещё одна деталь реализации — вес RMSNorm хранится как добавка к единице, .
Научный вклад
Заголовок раздела «Научный вклад»Gemma Team (2024) выпустили семейство открытых моделей, построенных, по словам авторов, на исследованиях и технологиях, созданных для Gemini:
- Два размера — 2B и 7B, для каждого предобученный и инструктивный (instruction-tuned) чекпоинт. 2B обучена на 3 трлн токенов, 7B — на 6 трлн, преимущественно английских текстов: веб-документы, математика, код.
- Токенизатор — подмножество SentencePiece-токенизатора Gemini со словарём 256k: цифры разбиваются по одной, лишние пробелы сохраняются, неизвестные символы кодируются байтами (tokenization.md).
- Архитектура (разд. 2 статьи) — decoder-only трансформер с контекстом 8192 токена и набором известных улучшений:
- Multi-Query Attention в 2B; в 7B — обычный multi-head attention (выбор MQA для 2B авторы обосновывают абляциями, по которым MQA хорошо работает на малом масштабе);
- RoPE в каждом слое (positional-encoding.md);
- GeGLU вместо ReLU (feed-forward.md);
- RMSNorm (normalization.md);
- общие эмбеддинги входа и выхода (weight tying, embeddings.md) — при словаре 256k это экономит сотни миллионов параметров.
- Качество. По заявлению авторов, Gemma превосходит открытые модели сопоставимого размера на 11 из 18 текстовых задач.
Гиперпараметры (табл. 1 статьи):
| Gemma 2B | Gemma 7B | |
|---|---|---|
| (d_model) | 2048 | 3072 |
| (слоёв) | 18 | 28 |
| feedforward hidden dims | 32768 | 49152 |
| (голов Q) | 8 | 16 |
| (голов K/V) | 1 | 16 |
| (размер головы) | 256 | 256 |
| словарь | 256128 | 256128 |
Два места этой таблицы нужно читать внимательно:
- «Feedforward hidden dims» — сумма размеров ветвей
gateиupGeGLU. Каждая из них имеет ширину (2B) и (7B) — так в HF (intermediate_size) и в этом репозитории. Это . - В 7B : проекции Q/K/V расширяют пространство, а сжимает его обратно. Поэтому в конфиге нужен явный
head_size.
В статье не упоминаются две детали, которые есть в эталонном коде (gemma_pytorch, HF GemmaModel): умножение эмбеддингов на и RMSNorm с множителем . Без них веса Gemma не воспроизводятся. Кроме того, в статье сказано, что нормализуется и вход, и выход каждого подслоя; в опубликованной реализации HF Gemma (первого поколения) — только вход (pre-LN, input_layernorm и post_attention_layernorm перед FFN). Этот репозиторий повторяет реализацию HF.
Архитектура блока декодера
Заголовок раздела «Архитектура блока декодера»%%{init: {"flowchart": {"rankSpacing": 28, "nodeSpacing": 28}}}%%
flowchart TB
Ids(["token ids"]):::io --> TokEmb["Token Embedding<br/>× √d, если scale_embeddings"]:::blue
TokEmb --> Drop["Dropout"]:::gray
subgraph Dec["GemmaDecoder × num_layers · pre-RMSNorm"]
direction TB
X(["x"]):::io --> N1["RMSNorm"]:::gray
N1 --> Attn["Grouped Query Attention<br/>num_kv_heads K/V-голов (1 — MQA)"]:::blueHl
R["RoPE<br/>cos/sin от позиции · без параметров<br/>один модуль на все слои"]:::rope
R -. "поворот Q и K" .-> Attn
Attn --> A1(("+")):::add
X -. residual .-> A1
A1 --> N2["RMSNorm"]:::gray
N2 --> FFN["GeGLU"]:::purpleHl
FFN --> A2(("+")):::add
A1 -. residual .-> A2
end
Drop --> Dec
Dec --> NF["RMSNorm<br/>(финальный)"]:::gray --> Lin
Lin["Linear → vocab_size<br/>(с tie_word_embeddings — матрица эмбеддингов)"]:::gray --> Out(["logits"]):::io
Out -. "generate(): softmax → выбор токена" .-> Next(["следующий токен"]):::io
style Dec fill:transparent,stroke:#82b366,stroke-width:2px,color:#5b9a3c
classDef io fill:#ffffff,stroke:#999999,color:#1a1a1a;
classDef add fill:#ffffff,stroke:#666666,color:#1a1a1a;
classDef blue fill:#dae8fc,stroke:#6c8ebf,color:#1a1a1a;
classDef blueHl fill:#dae8fc,stroke:#2f5f9e,stroke-width:3px,color:#1a1a1a;
classDef purple fill:#e1d5e7,stroke:#9673a6,color:#1a1a1a;
classDef purpleHl fill:#e1d5e7,stroke:#6a3d85,stroke-width:3px,color:#1a1a1a;
classDef gray fill:#f5f5f5,stroke:#666666,color:#1a1a1a;
classDef grayHl fill:#f5f5f5,stroke:#333333,stroke-width:3px,color:#1a1a1a;
classDef gold fill:#fff2cc,stroke:#d6b656,color:#1a1a1a;
classDef rope fill:#d5f0ec,stroke:#3a9e8f,color:#1a1a1a;
classDef ropeHl fill:#d5f0ec,stroke:#1f6f63,stroke-width:3px,color:#1a1a1a;
classDef dim fill:#f5f5f5,stroke:#bbbbbb,color:#999999,stroke-dasharray:4 3;
Как RoPE поворачивает Q и K — в разделе Attention с RoPE документа LLaMA.
Прямой проход в формулах
Заголовок раздела «Прямой проход в формулах»Формулы — для Gemma «как в статье» (все ключи из таблицы включены). Вход — индексы токенов ; — скрытые состояния после блока .
Эмбеддинги с масштабом:
где — матрица эмбеддингов, — её строки для токенов входа (строка — ), — вход первого блока (dropout в формулах опущен: «как в статье» он равен 0), — константа ( для 2B, для 7B). Тот же множитель был в исходном трансформере (Vaswani et al., 2017, разд. 3.4). Зачем он при связанных весах — в embeddings.md (раздел «Масштабирование эмбеддингов на √d»): матрица подобрана под роль выходной проекции, и её строки малы для входа в residual-поток.
RMSNorm Gemma для строки :
где — обучаемый вес в параметризации Gemma (инициализируется нулями), . При множитель равен 1 — то же, что вес RMSNorm LLaMA, инициализированный единицами. Функционально это одна и та же нормализация: . Отличается лишь то, что хранится в чекпоинте, — и, если применять weight decay к весам нормализации, к чему он их тянет: к 1 в параметризации Gemma, к 0 в обычной.
Пример: , : , нормализованный вектор . С множитель , результат . В этом репозитории тот же результат даёт RMSNorm с весом .
Блок (pre-LN, как у LLaMA и Mistral):
где — состояние после attention-подслоя блока , — две нормализации блока со своими весами.
Attention — GQA с RoPE (attention.md): для нормализованного входа и головы Q с номером группы K/V (как в mixtral.md)
где — выход головы , — проекции головы Q и группы K/V , , — causal-маска (masks.md), скользящего окна нет. У 2B : все 8 голов Q смотрят в одну пару K/V (MQA). У 7B (MHA) и , .
GeGLU (feed-forward.md) для каждой строки матрицы :
где , , . От SwiGLU отличается только активацией на ветви gate: вместо SiLU. Для примера: , (у SiLU — 0.7311 и −0.2689).
Выход — через ту же матрицу :
где — логиты, — финальная нормализация. Логит токена — скалярное произведение нормализованного скрытого состояния со строкой (embeddings.md). Bias нет нигде.
Multi-Query Attention vs GQA
Заголовок раздела «Multi-Query Attention vs GQA»MQA предложена в Shazeer, 2019, GQA — в Ainslie et al., 2023 как обобщение между MQA и MHA. Gemma 2B использует MQA (одна K/V-голова), Gemma 7B — обычный MHA (16 K/V-голов, по одной на Q-голову). Поэтому блок Gemma строится на GroupedQueryAttention (core/group_query_attention.py) без скользящего окна с num_kv_heads из конфига: 1 (по умолчанию) — MQA, num_q_heads — MHA. При одной K/V-голове она не копируется на все Q-головы, а транслируется в матричном умножении, так что результат побитово совпадает с прежним MultiQueryAttention (core/multi_query_attention.py); тот остался в llm.core как отдельный учебный модуль. KV-кэш слоя — тройка (K, V, next_pos), как у Mistral. Сравнение MHA, GQA и MQA — в attention.md.
Выигрыш MQA — в размере KV-кэша: на токен и слой он хранит чисел. У Gemma 2B это , а при MHA с теми же головами было бы — в 8 раз больше. На 18 слоях и контексте 8192 токена в bfloat16 (2 байта) — байт МиБ против МиБ ГиБ на одну последовательность (как в таблице attention.md).
Компоненты
Заголовок раздела «Компоненты»| Компонент | Класс | Файл |
|---|---|---|
| Токен-эмбеддинги | TokenEmbeddings | core/token_embeddings.py |
| Выходная проекция | nn.Linear или output_projection(..., tie_weights=True) | core/token_embeddings.py |
| Позиционное кодирование | RoPE | core/rope.py |
| Нормализация | RMSNorm | core/rms_norm.py |
| Attention | GroupedQueryAttention (num_kv_heads K/V-голов, по умолчанию 1 — MQA; RoPE; без окна) | core/group_query_attention.py |
| FFN | GeGLU (gated GELU-MLP) | core/geglu.py |
| Блок декодера | GemmaDecoder (pre-LN) | core/gemma_decoder.py |
| Модель целиком | Gemma | models/gemma/gemma.py |
| Перенос весов HF | convert_hf_state_dict | models/gemma/hf_weights.py |
Разбор кода
Заголовок раздела «Разбор кода»GemmaDecoder
Заголовок раздела «GemmaDecoder»core/gemma_decoder.py, GemmaDecoder(num_q_heads, emb_size, head_size, max_seq_len, rope, dropout=0.1, norm_eps=1e-6, num_kv_heads=1, intermediate_size=None, bias=True):
| Атрибут | Модуль | Формула |
|---|---|---|
_heads | GroupedQueryAttention(num_q_heads, num_kv_heads, emb_size, head_size, max_seq_len, rope, dropout, bias) — без window_size | |
_ff | GeGLU(emb_size, dropout, hidden_dim=intermediate_size, bias) | |
_norm1, _norm2 | RMSNorm(emb_size, eps=norm_eps) | , |
forward(x, use_cache=True, cache=None) — формулы блока:
norm1_out = self._norm1(x)attention, kv_caches = self._heads(norm1_out, use_cache=use_cache, cache=cache)out = attention + x # U^(l) = H^(l-1) + GQA(RMSNorm1(H^(l-1)))norm2_out = self._norm2(out)ffn_out = self._ff(norm2_out) # GeGLU(RMSNorm2(U^(l)))# возвращает (ffn_out + out, kv_caches) при use_cache, иначе (ffn_out + out, None)GeGLU (core/geglu.py) — три nn.Linear (_gate, _up, _down) и активация GELU из core/gelu.py — tanh-аппроксимация, совпадающая с gelu_pytorch_tanh в HF: out = self._down(self._up(x) * GELU(self._gate(x))), затем dropout.
RMSNorm (core/rms_norm.py) хранит вес _w, инициализированный единицами, и возвращает self._w * norm_x — то есть параметр , а не . Для float16/bfloat16 нормализация считается во float32, результат приводится к dtype входа и только потом умножается на вес.
models/gemma/gemma.py, наследник BaseModel.
__init__(config):
head_size = resolve_head_size(config, "num_q_heads", rope=True)—head_sizeиз конфига илиembed_dim // num_q_heads;- необязательные ключи:
rms_norm_eps(1e-6),intermediate_size(None→4 · embed_dimвнутриGeGLU),bias(True),num_kv_heads(1),rope_theta(10000),scale_embeddings(False),tie_word_embeddings(False); self._embedding_scale = math.sqrt(config["embed_dim"]), еслиscale_embeddings, иначеNone;_token_embeddings, один_position_embeddings(RoPE) на все слои,_dropout,_decodersизnum_layersблоковGemmaDecoder, финальная_norm;- выходная проекция: при
tie_word_embeddings—output_projection(self._token_embeddings, tie_weights=True):nn.Linearбез bias, чейweight— тот же объектnn.Parameter, что и матрица эмбеддингов; иначе — отдельныйnn.Linear(embed_dim, vocab_size, bias=bias); - инициализация как в HF:
init_normal_—LinearиEmbeddingиз (ключinitializer_range), bias — нули.
forward(x, use_cache=False, cache=None, attention_mask=None):
start_pos = cache_start_pos(cache)check_sequence_length(x.size(1), start_pos, self._max_seq_len)padding = padding_from_attention_mask(attention_mask, x, start_pos) # паддинг где угодно; см. masks.mdtok_out = self._token_embeddings(x) # [B, T, d]if self._embedding_scale is not None: tok_out = tok_out * torch.tensor(self._embedding_scale, dtype=tok_out.dtype) # × √dout = self._dropout(tok_out)# ... цикл по self._decoders с кэшем, как у Mixtral ...logits = self._linear(self._norm(out)) # Z = RMSNorm_f(H^(L)) E^T при tyingМножитель приводится к dtype эмбеддингов до умножения, как в HF: в bfloat16 округляется до 45.25, а — до 55.5. Это важно для побитового совпадения с HF в половинной точности.
convert_hf_state_dict
Заголовок раздела «convert_hf_state_dict»def convert_hf_state_dict(hf_state_dict: dict, num_heads: int, num_kv_heads: int = None) -> dict: result = _convert_llama_family(hf_state_dict, num_heads=num_heads, num_kv_heads=num_kv_heads) for key in result: if key.endswith("._w"): # веса RMSNorm: (1 + w) в HF → w здесь result[key] = result[key] + 1 return resultДва шага:
- Перенос LLaMA (
convert_hf_state_dictизmodels/llama/hf_weights.py): переименованиеmodel.embed_tokens→_token_embeddings._embedding,self_attn.{q,k,v,o}_proj→_heads._{q,k,v,layer},mlp.{gate,up,down}_proj→_ff._{gate,up,down},input_layernorm/post_attention_layernorm→_norm1/_norm2,model.norm→_norm,lm_head→_linear. Строкиq_projиk_projпереставляются внутри каждой головы: HF хранит пары RoPE как «первая половина головы | вторая половина», аRoPEздесь вращает соседние координаты (llama.md). Поэтому нужныnum_headsиnum_kv_heads. Еслиlm_head.weightв чекпоинте нет,_linear.weightполучает копию эмбеддингов. - Поправка RMSNorm: ко всем весам
._w(обе нормы каждого блока и финальная) прибавляется 1 — переход от к .
При связанных весах state_dict модели содержит и _token_embeddings._embedding.weight, и _linear.weight — это один параметр под двумя именами; load_state_dict записывает в него одно и то же значение дважды.
Подсчёт параметров
Заголовок раздела «Подсчёт параметров»С weight tying и без bias:
Неэмбеддинговые (non-embedding) параметры — всё, кроме .
Gemma 2B: , , , , , , .
| Часть | Формула | Параметров |
|---|---|---|
| attention слоя | 9 437 184 | |
| GeGLU слоя | 100 663 296 | |
| RMSNorm слоя | 4 096 | |
| слой | 110 104 576 | |
| 18 слоёв + финальная норма | 1 981 884 416 | |
| эмбеддинги | 524 288 000 | |
| всего | 2 506 172 416 ≈ 2.5 млрд |
Gemma 7B: , , , , .
| Часть | Формула | Параметров |
|---|---|---|
| attention слоя | 50 331 648 | |
| GeGLU слоя | 226 492 416 | |
| RMSNorm слоя | 6 144 | |
| слой | 276 830 208 | |
| 28 слоёв + финальная норма | 7 751 248 896 | |
| эмбеддинги | 786 432 000 | |
| всего | 8 537 680 896 ≈ 8.5 млрд |
Сравнение со статьёй (табл. 2):
| эмбеддинговые, статья | эмбеддинговые, здесь | неэмбеддинговые, статья | неэмбеддинговые, здесь | |
|---|---|---|---|---|
| 2B | 524 550 144 | 524 288 000 | 1 981 884 416 | 1 981 884 416 |
| 7B | 786 825 216 | 786 432 000 | 7 751 248 896 | 7 751 248 896 |
Неэмбеддинговые параметры совпадают до единицы — это подтверждает, что архитектура (размеры, отсутствие bias, по две нормы на блок и финальная) воспроизведена точно. Эмбеддинговые в статье посчитаны для словаря 256 128 строк (, ), а в конфиге HF vocab_size = 256000; разница — 128 строк.
Заметьте, что «7B» — это неэмбеддинговые 7.75 млрд; с эмбеддингами модель содержит 8.5 млрд. Эмбеддинги у 2B — 21% всех параметров: без weight tying отдельная выходная матрица добавила бы ещё 524 млн (всего 3.03 млрд).
Проверка на meta-устройстве (память под веса не выделяется):
import torchfrom llm.models.gemma import Gemma
cfg = {"vocab_size": 256000, "embed_dim": 2048, "num_q_heads": 8, "num_kv_heads": 1, "head_size": 256, "num_layers": 18, "max_position_embeddings": 8192, "dropout": 0.0, "intermediate_size": 16384, "bias": False, "tie_word_embeddings": True, "scale_embeddings": True}with torch.device("meta"): model = Gemma(cfg)total = sum(p.numel() for p in model.parameters()) # общий параметр считается один разemb = model._token_embeddings._embedding.weight.numel()print(total, emb, total - emb) # 2506172416 524288000 1981884416model.parameters() не повторяет один и тот же nn.Parameter, поэтому связанная матрица учтена один раз.
Учебный конфиг gemma_train.json: , , , (по умолчанию), , , , bias, отдельная выходная проекция:
| Часть | Параметров |
|---|---|
| attention слоя: | 164 480 |
| GeGLU слоя: | 788 736 |
| слой (с двумя нормами по 256) | 953 728 |
| эмбеддинги + выход с bias + финальная норма | 513 256 |
| всего | 4 328 168 |
Отличия от Gemma
Заголовок раздела «Отличия от Gemma»Сравнение с Gemma 2B/7B (статья и GemmaConfig/GemmaModel в HF). Подробности, воспроизведение и варианты исправления — в бэклоге (номера пунктов в скобках).
| Gemma | Здесь | |
|---|---|---|
| Масштаб эмбеддингов | умножаются на √d перед первым блоком | по умолчанию нет; scale_embeddings: true — как в оригинале (42) |
| Выходная проекция | привязана к эмбеддингам (tie_word_embeddings) | по умолчанию отдельный Linear; tie_word_embeddings: true — как в оригинале (43) |
| Bias | нет ни в одной проекции | по умолчанию во всех Linear; bias: false — как в оригинале (43) |
| Скрытый слой GeGLU | 8·d на каждую из gate/up (16384 при d = 2048) | по умолчанию 4·d; intermediate_size — любой (44) |
| Attention | 2B — MQA, 7B — MHA с 16 головами и head_dim = 256 ≠ d / heads | по умолчанию MQA; num_kv_heads и head_size из конфига (45) |
| RMSNorm | вес с нуля, множитель (1 + w), вычисление во float32 | вес с единиц, множитель w — при загрузке весов HF к ним прибавляется 1; для float16/bfloat16 нормализация во float32 (46) |
| Dropout | нет | после эмбеддингов, в attention и в GeGLU (55); dropout: 0 убирает его полностью |
Активация GeGLU — tanh-аппроксимация GELU — совпадает с оригиналом (gelu_pytorch_tanh в HF).
Конфигурация
Заголовок раздела «Конфигурация»Пример из experiments/llm_only/configs/gemma_train.json:
| Параметр | Значение в примере | Используется? |
|---|---|---|
vocab_size | (из токенизатора) | ✅ |
embed_dim | 256 | ✅ |
num_q_heads | 4 | ✅ число Query-голов; число K/V-голов — необязательный num_kv_heads (по умолчанию 1 — MQA, см. Как в статье) |
num_layers | 4 | ✅ |
max_position_embeddings | 512 | ✅ |
rms_norm_eps | (нет в примере) | ✅ необязательный eps всех RMSNorm, по умолчанию 1e-6 — как в Gemma |
rope_theta | (нет в примере) | ✅ необязательная база частот RoPE, по умолчанию 10000 — как в Gemma; см. llama.md |
initializer_range | (нет в примере) | ✅ необязательное стандартное отклонение начальных весов Linear и Embedding, по умолчанию 0.02 — как в HF; см. training.md |
dropout | 0.1 | ✅ после эмбеддингов, в attention и GeGLU; в Gemma dropout нет — для соответствия оригиналу 0 |
head_size | 64 | ✅ необязательный; по умолчанию embed_dim // num_q_heads |
Как в статье
Заголовок раздела «Как в статье»Необязательные ключи; без них структура модели прежняя, и старые чекпоинты загружаются. Все, кроме scale_embeddings, меняют форму весов, поэтому чекпоинт одного вида в модель другого не загрузится.
| Ключ | По умолчанию | Gemma 2B | Gemma 7B |
|---|---|---|---|
num_kv_heads | 1 (MQA) | 1 | 16 |
head_size | embed_dim // num_q_heads | 256 | 256 (≠ 3072 / 16) |
intermediate_size | 4 · embed_dim | 16384 (8·d) | 24576 (8·d) |
bias | true | false | false |
tie_word_embeddings | false | true | true |
scale_embeddings | false | true | true |
rms_norm_eps | 1e-6 | 1e-6 | 1e-6 |
dropout | — | 0 | 0 |
scale_embeddings умножает выход эмбеддингов на √embed_dim (множитель приводится к dtype эмбеддингов, как в HF). При tied embeddings одна матрица служит и входом, и выходом, и её норма рассчитана на выходную проекцию; без множителя вход в первый блок был бы на порядок меньше. tie_word_embeddings особенно заметен у Gemma: словарь 256 000 токенов, и отдельная голова для 2B — это ещё ~524M параметров.
При обучении с нуля важна инициализация. Gemma, как HF, инициализирует Linear и Embedding из (init_normal_, ключ initializer_range), и с tie_word_embeddings и scale_embeddings начальный cross-entropy на учебном конфиге (, ) — 7.06 при (std логитов 0.36). С инициализацией nn.Embedding по умолчанию, , те же ключи дают логиты порядка : стандартное отклонение ≈18, начальный cross-entropy ≈258. Почему нужны малые эмбеддинги — в embeddings.md (раздел «Масштабирование эмбеддингов на √d»).
Загрузка весов HuggingFace
Заголовок раздела «Загрузка весов HuggingFace»С ключами из таблицы выше загружаются веса GemmaForCausalLM — через convert_hf_state_dict из models/gemma/hf_weights.py. Это перенос LLaMA (llama.md: те же имена слоёв и перестановка строк q_proj/k_proj под RoPE на чередующихся парах) плюс одна поправка: GemmaRMSNorm умножает на (1 + w), а RMSNorm здесь — на w, поэтому к весам всех RMSNorm прибавляется 1 (разбор — в Разборе кода).
from transformers import GemmaForCausalLMfrom llm.models.gemma import Gemma, convert_hf_state_dict
hf = GemmaForCausalLM.from_pretrained("google/gemma-2b")c = hf.configmodel = Gemma({"vocab_size": c.vocab_size, "embed_dim": c.hidden_size, "num_q_heads": c.num_attention_heads, "num_kv_heads": c.num_key_value_heads, "head_size": c.head_dim, "num_layers": c.num_hidden_layers, "max_position_embeddings": c.max_position_embeddings, "dropout": 0.0, "rms_norm_eps": c.rms_norm_eps, "rope_theta": c.rope_theta, "intermediate_size": c.intermediate_size, "bias": False, "tie_word_embeddings": True, "scale_embeddings": True})model.load_state_dict(convert_hf_state_dict(hf.state_dict(), num_heads=c.num_attention_heads, num_kv_heads=c.num_key_value_heads))Сверено со случайными GemmaForCausalLM из transformers в двух формах — MQA с head_dim = hidden / heads (как 2B) и MHA с head_dim ≠ hidden / heads (как 7B): логиты совпадают до ~1e-5, greedy-генерация с KV-кэшем — токен в токен (llm/tests/models/test_gemma_hf_parity.py). Без scale_embeddings или без +1 к весам RMSNorm результат HF не воспроизводится — это тоже проверяет тест. Настоящие веса google/gemma-2b закрыты лицензией (доступ после принятия условий на HuggingFace), в проверке они не использовались.
В bfloat16 возможна разница в последних битах: GemmaRMSNorm умножает на вес ещё во float32, а RMSNorm здесь — после приведения к dtype входа, как LlamaRMSNorm.
Генерация
Заголовок раздела «Генерация»Gemma.generate(...) — унифицированная сигнатура (см. gpt.md и generation.md). Благодаря MQA KV-кэш Gemma 2B в 8 раз меньше, чем при MHA с тем же числом голов Q (см. выше).
Итоги линейки
Заголовок раздела «Итоги линейки»Шесть моделей части II — это одна и та же схема decoder-only трансформера (language-modeling.md), в которой менялись отдельные узлы:
| позиция | нормализация | attention | FFN | выход | |
|---|---|---|---|---|---|
| GPT-1 | обучаемые эмбеддинги | LayerNorm, post-LN | MHA | GELU | связан с эмбеддингами (опция) |
| GPT-2 | обучаемые эмбеддинги | LayerNorm, pre-LN | MHA | GELU | связан с эмбеддингами (опция) |
| LLaMA | RoPE | RMSNorm, pre-LN | MHA | SwiGLU | отдельный |
| Mistral | RoPE | RMSNorm, pre-LN | GQA + окно | SwiGLU | отдельный |
| Mixtral | RoPE | RMSNorm, pre-LN | GQA | MoE из SwiGLU | отдельный |
| Gemma | RoPE | RMSNorm , pre-LN | MQA / MHA | GeGLU | связан, эмбеддинги × √d |
- Gemma — decoder-only трансформер на базе RoPE + RMSNorm с GeGLU ( на каждую ветвь), MQA в 2B и MHA в 7B, словарём 256k, связанными эмбеддингами и выходом и умножением эмбеддингов на .
- У 7B , поэтому в конфиге нужен явный
head_size. - RMSNorm Gemma хранит добавку к единице; при загрузке весов HF к весам нормализаций прибавляется 1.
- Параметры: 2 506 172 416 у 2B и 8 537 680 896 у 7B при словаре 256 000; неэмбеддинговые совпадают с табл. 2 статьи до единицы.
- В библиотеке оригинальная структура включается ключами
num_kv_heads,head_size,intermediate_size,bias: false,tie_word_embeddings,scale_embeddings; по умолчанию они выключены ради старых чекпоинтов.
Вопросы и упражнения
Заголовок раздела «Вопросы и упражнения»-
Почему в Gemma 7B — и какие матрицы из-за этого не квадратные? Запишите их формы.
Ответ
, . (в
nn.Linear—weightформы[4096, 3072]), . Attention работает в пространстве размерности 4096 и сжимает результат обратно в 3072. -
Сколько параметров сэкономила бы Gemma 2B, если бы вместо MQA использовала GQA с , — или, наоборот, сколько бы добавилось? А при MHA ()?
Ответ
K и V слоя: . При — 1 048 576, при — 2 097 152 (+1 048 576 на слой, +18 874 368 на модель), при — 8 388 608 (+7 340 032 на слой, +132 120 576 ≈ 132 млн на модель). MQA экономит немного параметров, но главное — в 8 раз меньший KV-кэш.
-
Вес RMSNorm в чекпоинте HF равен . Какое значение окажется в
_wпослеconvert_hf_state_dict? Каким будет выход для нормализованной координаты 2.0?Ответ
_w = 1 + (−0.3) = 0.7; выход — так же, как в HF: . -
Токен имеет эмбеддинг с RMS координат 0.02 (инициализация HF). Каков RMS после
scale_embeddingsу Gemma 2B и 7B?Ответ
и — порядка 1, как и задумано.
-
Повторите подсчёт параметров Gemma 7B на
meta-устройстве. Затем выключитеtie_word_embeddings. Сколько параметров добавится и почему у выходной проекции не появляется bias приbias: false?Ответ
Всего с tying — 8 537 680 896; без — 9 324 112 896, то есть . Отдельная проекция создаётся как
nn.Linear(embed_dim, vocab_size, bias=bias), и приbias: falsebias у неё нет. (При tying bias нет в любом случае:output_projection(..., tie_weights=True)создаётLinearсbias=not tie_weights.) -
Что сломается, если загрузить веса HF Gemma без
+1к весам RMSNorm? Оцените: чему равен множитель после нормализации у только что инициализированной HF-модели, если её веса загрузить без поправки?Ответ
У свежей HF-модели . Без поправки
_w = 0, и каждая RMSNorm выдаёт нули. Attention и GeGLU без bias на нулевом входе дают нули, так что каждый блок просто пропускает residual-поток (эмбеддинги) без изменений; финальная норма обнуляет и его, и все логиты равны 0. У обученных весов не нули, но результат всё равно неверен — множители на 1 меньше нужных. -
(Код.) В
Gemma.forwardмножитель сначала превращается в тензор с dtype эмбеддингов. Проверьте в Python, какое значение получается для в bfloat16, и объясните, почемуtok_out * math.sqrt(3072)дал бы в bfloat16 другой результат, чем HF.Ответ
torch.tensor(3072 ** 0.5, dtype=torch.bfloat16)— 55.5 (у bfloat16 8 бит мантиссы, соседние представимые числа около 55 отстоят на 0.25). HF умножает на это округлённое значение. При умножении bfloat16-тензора на Python-float PyTorch использует неокруглённый множитель 55.4256… и округляет только произведение, поэтому результат другой: на 10 000 случайных bfloat16-числах он отличается от умножения на 55.5 примерно в четверти элементов. Явный тензор с dtype эмбеддингов даёт тот же множитель, что и в HF. -
(Ноутбук.) В
notebooks/gemma.ipynbобучите учебную Gemma дважды: с ключами по умолчанию и сtie_word_embeddings: true,scale_embeddings: true. Сравните начальный loss и кривые обучения. Как исправить начальный loss во втором случае, не меняя код модели? (Подсказка:llm.core.weight_init.)
Литература
Заголовок раздела «Литература»Основная статья:
- Gemma Team. Gemma: Open Models Based on Gemini Research and Technology. 2024. arXiv:2403.08295
Компоненты:
- Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019. arXiv:1911.02150 — Multi-Query Attention
- Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. 2023. arXiv:2305.13245
- Shazeer. GLU Variants Improve Transformer. 2020. arXiv:2002.05202 — SwiGLU и GeGLU
- Hendrycks, Gimpel. Gaussian Error Linear Units (GELUs). 2016. arXiv:1606.08415
- Su et al. RoFormer: Enhanced Transformer with Rotary Position Embedding. 2021. arXiv:2104.09864
- Zhang, Sennrich. Root Mean Square Layer Normalization. 2019. arXiv:1910.07467
- Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762 — масштабирование эмбеддингов на √d (разд. 3.4)