Mistral
Реализация:
llm/src/llm/models/mistral/mistral.py· классMistralНоутбук:notebooks/mistral.ipynb
Место в линейке: GPT-1 → GPT-2 → LLaMA → Mistral → Mixtral · Gemma
Что вы узнаете
Заголовок раздела «Что вы узнаете»- Что предложила статья Mistral 7B и за счёт чего модель на 7B обходит Llama 2 13B.
- Как устроен Grouped Query Attention и сколько памяти и параметров он экономит.
- Как скользящее окно (sliding window attention) ограничивает внимание, почему дальность зависимостей всё равно растёт с глубиной и почему окно здесь шириной .
- Как работает кольцевой кэш (rolling buffer cache), чем его заменяет эта реализация и как обрабатывать длинный промпт кусками (chunked prefill).
- Как посчитать 7,24 млрд параметров Mistral 7B и проверить подсчёт программно.
- Что в классах
Mistral,MistralDecoderиGroupedQueryAttentionотличается от LLaMA.
Предварительные знания
Заголовок раздела «Предварительные знания»- Архитектура LLaMA: RMSNorm, SwiGLU, RoPE — llama.md.
- Multi-head attention, MHA/GQA/MQA, KV-кэш — attention.md.
- Causal-маска и маска скользящего окна — masks.md.
- Авторегрессивная генерация с кэшем — generation.md.
Mistral 7B (Jiang et al., Mistral 7B, Mistral AI, 2023) — decoder-only модель на 7,24 млрд параметров. Её блок — это блок LLaMA (pre-RMSNorm, SwiGLU, RoPE) с двумя изменениями в attention, нацеленными на дешёвый инференс:
- Grouped Query Attention (GQA, Ainslie et al., 2023): 32 головы Q делят 8 голов K/V — KV-кэш в 4 раза меньше;
- Sliding Window Attention (SWA): каждый токен в слое смотрит только на последние позиций — кэш можно ограничить окном.
Научный вклад
Заголовок раздела «Научный вклад»Статья показывает, что аккуратно спроектированная маленькая модель может превзойти модели крупнее (аннотация и разд. 3 статьи):
- Mistral 7B превосходит Llama 2 13B на всех бенчмарках, которые оценивали авторы, и LLaMA 1 34B (так в статье Mistral; у Meta эта модель — 33B) — на задачах рассуждения, математики и генерации кода;
- дообученная для диалога Mistral 7B – Instruct превосходит Llama 2 13B – Chat.
Архитектурные средства (разд. 2):
- GQA ускоряет инференс и уменьшает память кэша, позволяя увеличить батч;
- SWA снижает стоимость внимания на длинных последовательностях: для длины 16K и доработки FlashAttention и xFormers дают ускорение в 2 раза по сравнению с обычным вниманием;
- rolling buffer cache — кэш фиксированного размера с записью по позиции : на последовательности 32K память кэша уменьшается в 8 раз без потери качества;
- pre-fill и chunking — промпт известен заранее, поэтому кэш заполняется им сразу, а длинный промпт — кусками размером с окно.
Веса опубликованы под лицензией Apache 2.0.
Гиперпараметры (табл. 1 статьи):
| Параметр | Значение | Здесь |
|---|---|---|
dim | 4096 | embed_dim () |
n_layers | 32 | num_layers () |
head_dim | 128 | head_size () |
hidden_dim | 14336 | intermediate_size () |
n_heads | 32 | num_q_heads () |
n_kv_heads | 8 | num_kv_heads () |
window_size | 4096 | window_size (, но см. Ширина окна: W + 1) |
context_len | 8192 | max_position_embeddings () |
vocab_size | 32000 | vocab_size () |
Архитектура блока декодера
Заголовок раздела «Архитектура блока декодера»Жирная обводка — то, что изменилось по сравнению с LLaMA.
%%{init: {"flowchart": {"rankSpacing": 28, "nodeSpacing": 28}}}%%
flowchart TB
Ids(["token ids"]):::io --> TokEmb["Token Embedding"]:::blue
TokEmb --> Drop["Dropout"]:::gray
subgraph Dec["MistralDecoder × num_layers · pre-RMSNorm"]
direction TB
X(["x"]):::io --> N1["RMSNorm"]:::gray
N1 --> Attn["Grouped Query Attention<br/>sliding window"]:::blueHl
R["RoPE<br/>cos/sin от позиции · без параметров<br/>один модуль на все слои"]:::rope
R -. "поворот Q и K" .-> Attn
Attn --> A1(("+")):::add
X -. residual .-> A1
A1 --> N2["RMSNorm"]:::gray
N2 --> FFN["SwiGLU"]:::purple
FFN --> A2(("+")):::add
A1 -. residual .-> A2
end
Drop --> Dec
Dec --> NF["RMSNorm<br/>(финальный)"]:::gray --> Lin
Lin["Linear → vocab_size"]:::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;
Прямой проход тот же, что у LLaMA (формулы — в llama.md), с заменой на с маской окна:
где — выход блока , — состояние после его attention-подслоя, — grouped query attention с RoPE и маской окна ширины (формулы ниже). Как RoPE поворачивает Q и K — в llama.md.
Grouped Query Attention
Заголовок раздела «Grouped Query Attention»Общая часть — виды attention по числу голов K/V и таблица размеров кэша реальных моделей — в attention.md. Здесь — формулы и то, что относится к Mistral.
В MHA у каждой из голов Q свои K и V. В GQA голов K/V меньше, , и делится на : головы Q разбиты на групп по , и группа делит одну пару K/V. Для головы :
где:
- — вход (выход RMSNorm);
- — проекция запроса головы , всего штук;
- — проекции ключа и значения группы , всего штук;
- — номер группы головы : головы — группа 0, следующие — группа 1 и т. д.;
- ;
- — маска скользящего окна (ниже);
- — выходная проекция.
— это MHA, — Multi-Query Attention (Shazeer, 2019). У Mistral 7B , , : головы Q 0–3 читают K/V группы 0, головы 4–7 — группы 1 и т. д.
Интуиция. Головы Q задают, что ищет токен, и их много. K и V — что токен предлагает; их разнообразие, как показали Ainslie et al., можно уменьшить в несколько раз почти без потери качества. Зато K и V — именно то, что хранится в кэше при генерации.
Сколько экономит GQA
Заголовок раздела «Сколько экономит GQA»KV-кэш. На токен в одном слое хранится чисел (K и V) вместо — в раз меньше. Для всей модели и последовательности длины :
где — байт на число (2 для float16/bfloat16), множитель 2 — K и V. Для Mistral 7B: байт КиБ на токен. При MHA () было бы 512 КиБ.
| Mistral 7B, float16, один запрос | 4096 токенов | 32 768 токенов |
|---|---|---|
| MHA (), полный кэш | 2 ГиБ | 16 ГиБ |
| GQA (), полный кэш | 512 МиБ | 4 ГиБ |
| GQA + кэш, ограниченный окном | 512 МиБ | 512 МиБ |
Последняя строка — вклад скользящего окна: кэш перестаёт расти после токенов; на 32K это те самые «в 8 раз» из статьи.
Параметры. и имеют форму вместо . Для Mistral 7B — M вместо M на слой; на 32 слоях экономия M параметров.
Вычисления. GQA не удешевляет: каждая из голов Q по-прежнему считает свои веса по всем ключам. Экономятся проекции K, V и — главное — чтение кэша из памяти, которое и ограничивает скорость генерации.
Sliding Window Attention
Заголовок раздела «Sliding Window Attention»Идея локального внимания со скользящим окном — из Longformer (Beltagy et al., 2020). В обычной causal-маске токен видит все позиции ; в маске окна — только последние:
где:
- — позиция запроса, — позиция ключа (абсолютные, с 0);
- —
window_size; - — пара разрешена, — после softmax её вес станет нулём.
Условие — обычная causal-часть, — окно. Токен видит позиций вместе с собой (почему именно столько — в Ширина окна: W + 1). Маска для , (строки — запросы, столбцы — ключи; вывод GroupedQueryAttention._tril_mask[:8, :8]):
j: 0 1 2 3 4 5 6 7i = 0: 1 . . . . . . .i = 1: 1 1 . . . . . .i = 2: 1 1 1 . . . . .i = 3: 1 1 1 1 . . . .i = 4: 1 1 1 1 1 . . .i = 5: . 1 1 1 1 1 . .i = 6: . . 1 1 1 1 1 .i = 7: . . . 1 1 1 1 1Пока , маска совпадает с causal; дальше диагональная полоса ширины сдвигается вправо.
Стоимость. Строка имеет не больше ненулевых весов, поэтому внимание на слой требует операций вместо — если ядро вычисляет только полосу. Реализация здесь этого не делает: GroupedQueryAttention считает полную матрицу и затем обнуляет лишнее маской. Экономия по памяти здесь есть только в кэше при генерации (ниже); по вычислениям при обучении и префилле её нет.
Рецептивное поле растёт со слоями
Заголовок раздела «Рецептивное поле растёт со слоями»Внутри одного слоя токен не видит дальше позиций назад. Но скрытое состояние позиции в слое уже содержит информацию о позициях до входа. По индукции (разд. 2 статьи):
- База: — это определение маски.
- Шаг: смотрит на с , а те — на входы с номерами .
После слоёв дальность — до токенов. Для Mistral 7B: K токенов — теоретическое поле внимания, в 16 раз больше контекста 8192. «Теоретическое» — потому что информация на каждом шаге передаётся через сжатое скрытое состояние и по пути теряется; это верхняя граница, а не гарантия.
Для учебного конфига (, ) поле — 64 токена при : модель видит в прошлое намного меньше своего контекста.
Ширина окна: W + 1
Заголовок раздела «Ширина окна: W + 1»Здесь W = window_size. Реализация пропускает W + 1 позиций: маска i − j ≤ W (включая сам токен). Источники определяют окно по-разному:
| Источник | Позиций видно (вместе с токеном) |
|---|---|
| Статья Mistral 7B, раздел 2, текст: «attends to all hidden states from the previous layer with positions between i − W and i» | W + 1 |
Эталонный код Mistral AI (one_file_ref.py), prefill: torch.triu(mask, diagonal=-sliding_window) | W + 1 |
| Статья, подпись к рисунку 1: «each token can attend to at most W tokens» | W |
| Статья, Rolling Buffer Cache: кэш фиксированного размера W, текущий токен тоже в нём | W |
| Эталонный код Mistral AI, генерация с кэшем (буфер из W ячеек) | W |
HuggingFace Transformers (sliding_window_overlay: kv_idx > q_idx − sliding_window) | W |
Реализация следует тексту статьи и prefill в эталонном коде, причём одинаково с кэшем и без. От HuggingFace она отличается на одну позицию: при загрузке весов Mistral из HF окно нужно задать на единицу меньше, window_size = sliding_window − 1 (см. Загрузка весов HuggingFace); с тем же числом логиты не совпадут — это проверяет тест. Для window_size = 4096 (Mistral 7B) разница несущественна, для учебных конфигов с window_size = 16 — около 6 %.
Кэш, ограниченный окном
Заголовок раздела «Кэш, ограниченный окном»Rolling buffer cache в статье
Заголовок раздела «Rolling buffer cache в статье»Ключи старше позиций больше никогда не понадобятся: ни один будущий запрос до них не дотянется. Поэтому кэш можно держать фиксированного размера и перезаписывать по кругу (разд. 2 статьи):
где — абсолютная позиция токена, ячейка — индекс в буфере из элементов (для каждого слоя и каждой головы K/V). K и V позиции записываются в ячейку и затирают то, что там лежало, — позицию .
Пример, , пишем позиции 0–9:
после позиции: ячейка 0 ячейка 1 ячейка 2 ячейка 3 3 0 1 2 3 4 4 1 2 3 ← 4 mod 4 = 0, затёрта позиция 0 5 4 5 2 3 9 8 9 6 7Порядок ячеек не совпадает с порядком позиций, но attention это не важно: softmax берётся по множеству ключей, а позиция уже «вшита» в K поворотом RoPE. Запись одного токена — без копирования.
Как это сделано здесь
Заголовок раздела «Как это сделано здесь»Кольцевого буфера здесь нет. GroupedQueryAttention.forward на каждом шаге приклеивает новые K/V к кэшу через torch.cat и обрезает результат срезом до последних позиций:
k = torch.cat([k_cache, k], dim=2) # [B, G, cache_len + T, d_h]...if self._window_size is not None: # после вычисления attention k = k[:, :, -self._window_size:, :] # остаются последние W позиций v = v[:, :, -self._window_size:, :]kv_cache = (k, v, start_pos + seq_len) # тройка: K, V, next_posСодержимое то же, что в кольцевом буфере, но в порядке позиций; цена — копирование на каждом шаге вместо записи .
Из-за обрезки длина кэша перестаёт совпадать с позицией токена, поэтому кэш — это тройка (K, V, next_pos): next_pos — абсолютная позиция следующего токена. Её RoPE использует как start_pos (start_pos = cache[2]), и по ней же модель проверяет длину (cache_start_pos в core/generation.py). Маска берётся срезом абсолютной маски:
cache_len = k.size(2) - seq_lenwindow_mask = self._tril_mask[start_pos : start_pos + seq_len, # строки — новые запросы start_pos - cache_len : start_pos + seq_len] # столбцы — ключи кэша и новыеПример, : после 6 токенов (позиции 0–5) кэш хранит K/V позиций 2, 3, 4, 5, next_pos = 6. Токен на позиции 6 видит кэш (2–5) и себя (6) — пять позиций, , как и маска без кэша: . В кольцевом буфере статьи тот же токен занял бы ячейку позиции 2 и видел бы позиции 3–6 — позиций (строка таблицы выше про рисунок и буфер).
Без window_size кэш не обрезается и растёт, как у LLaMA, но тоже остаётся тройкой.
Pre-fill и chunking
Заголовок раздела «Pre-fill и chunking»При генерации промпт известен целиком, поэтому его K и V можно вычислить за один проход — pre-fill — и только потом генерировать по токену. Если промпт очень длинный, матрица внимания не помещается в память; статья предлагает делить промпт на куски (chunking) размером с окно и заполнять кэш кусок за куском (разд. 2, рис. 3). Каждый кусок длины (в статье ) смотрит на кэш (предыдущее окно, не больше позиций) и на себя с causal-маской, поэтому матрица оценок внимания куска имеет размер не больше вместо .
Здесь. generate делает pre-fill всего промпта за один вызов forward. Префилл кусками получается вручную: forward с кэшем принимает кусок любой длины, и срез маски выше правильно обрабатывает несколько новых запросов сразу:
cache = Nonefor s in range(0, prompt.size(1), chunk): logits, cache = model(prompt[:, s:s + chunk], use_cache=True, cache=cache)next_token = logits[:, -1].argmax(-1, keepdim=True) # дальше — по токену с тем же кэшемПроверено: для модели с и 13 токенами логиты при кусках по 1, 3, 4 и 7 токенов совпадают с проходом целиком до ~5·10⁻⁷. Размер куска может быть и больше : кэш после куска всё равно обрезается до , а каждому запросу нужно не больше ключей до себя. Матрица внимания куска — [B, H, C, cache_len + C], где cache_len ≤ W.
Компоненты
Заголовок раздела «Компоненты»| Компонент | Класс | Файл |
|---|---|---|
| Токен-эмбеддинги | TokenEmbeddings | core/token_embeddings.py |
| Позиционное кодирование | RoPE | core/rope.py |
| Нормализация | RMSNorm | core/rms_norm.py |
| Attention | GroupedQueryAttention (GQA + скользящее окно + RoPE) | core/group_query_attention.py |
| FFN | SwiGLU | core/swi_glu.py |
| Блок декодера | MistralDecoder (pre-norm) | core/mistral_decoder.py |
| Модель целиком | Mistral | models/mistral/mistral.py |
Разбор кода
Заголовок раздела «Разбор кода»Общая часть — то же, что у LLaMA (Разбор кода): pre-norm блок, один RoPE на все слои, финальная RMSNorm, голова без связи с эмбеддингами, generate из BaseModel. Ниже — только отличия.
Класс Mistral
Заголовок раздела «Класс Mistral»- размер головы —
resolve_head_size(config, "num_q_heads", rope=True): ключhead_sizeилиembed_dim // num_q_heads; - в каждый
MistralDecoderпередаютсяnum_q_heads,num_kv_heads,window_size=config.get("window_size")(None— окна нет),norm_eps,intermediate_size(None— внутриSwiGLU) иbias; forwardустроен как уLlama; позиция для проверки длины берётся изnext_posкэша черезcache_start_pos— функция отличает тройку(K, V, next_pos)от пары(K, V)LLaMA.
Класс MistralDecoder
Заголовок раздела «Класс MistralDecoder»core/mistral_decoder.py. В отличие от параметризуемого CachedDecoder, здесь состав зафиксирован: GroupedQueryAttention (_heads), SwiGLU (_ff), две RMSNorm (_norm1, _norm2). Имена полей те же, что у CachedDecoder, поэтому одна функция convert_hf_state_dict обслуживает обе модели. forward:
norm1_out = RMSNorm1(x)attn_out = GQA(norm1_out) # с RoPE и маской окнаout = attn_out + xnorm2_out = RMSNorm2(out)ffn_out = SwiGLU(norm2_out)result = ffn_out + outКласс GroupedQueryAttention
Заголовок раздела «Класс GroupedQueryAttention»core/group_query_attention.py. Специфичное для GQA и окна:
Конструктор. Проверяет, что num_q_heads % num_kv_heads == 0 (иначе ValueError). Проекции: _q = nn.Linear(d, H·d_h), _k и _v = nn.Linear(d, G·d_h), _layer = nn.Linear(H·d_h, d) — это . Маска строится один раз на методом _create_sliding_window_mask:
causal_mask = col_indices <= row_indices # j ≤ iwindow_mask = row_indices - col_indices <= window_size # i − j ≤ Wmask = causal_mask & window_mask— ровно формула . Без window_size в качестве окна подставляется max_seq_len, и маска становится обычной causal.
Повтор голов K/V. После RoPE и склейки с кэшем K и V имеют форму [B, G, T_k, d_h], а Q — [B, H, T, d_h]. _repeat_kv_heads размножает K/V до голов:
kv = kv.unsqueeze(2) # [B, G, 1, T_k, d_h]kv = kv.repeat(1, 1, num_repeats, 1, 1) # [B, G, r, T_k, d_h]kv = kv.reshape(batch_size, num_q_heads, seq_len, head_size) # [B, H, T_k, d_h]После reshape голова с номером () получает копию группы — это и есть . Порядок важен для загрузки весов HF: там используется та же схема (repeat_kv). При повтор пропускается — [B, 1, T_k, d_h] транслируется (broadcast) по головам при умножении. Копии создаются только на время вычисления; в кэш идут K/V с головами.
Маска с кэшем берётся срезом по абсолютным позициям, кэш обрезается до и возвращается тройкой — см. Как это сделано здесь. Dropout — только на выходе ; на веса внимания не применяется.
Подсчёт параметров
Заголовок раздела «Подсчёт параметров»Обозначим , если bias: true, иначе . Отличие от LLaMA — только в K и V:
| Компонент | Параметров |
|---|---|
| Эмбеддинги | |
| , одного слоя | |
| , одного слоя | |
| SwiGLU одного слоя | |
| Две RMSNorm слоя | |
| Финальная RMSNorm | |
| Голова |
Итого без bias:
где — словарь, — embed_dim, — num_layers, , — головы Q и K/V, — head_size, — intermediate_size.
Mistral 7B: , , , , , , :
млрд. FFN — 81 % параметров слоя; attention — 19 % (у LLaMA 7B — 33 %): GQA урезал K/V, а FFN стал шире. С MHA () было бы 8 047 038 464.
Учебный конфиг experiments/llm_only/configs/mistral_train.json: , , , , , , , (из токенизатора):
Программная проверка на мета-устройстве (torch.device("meta"): тензоры без данных, память не выделяется; PyTorch ≥ 2.0):
import json, torchfrom llm.models.mistral import Mistral
def count(model): return sum(p.numel() for p in model.parameters())
cfg = json.load(open("experiments/llm_only/configs/mistral_train.json"))["model_config"]cfg["vocab_size"] = 1000print(count(Mistral(cfg))) # 4459752
cfg_7b = {"vocab_size": 32000, "embed_dim": 4096, "num_q_heads": 32, "num_kv_heads": 8, "head_size": 128, "num_layers": 32, "max_position_embeddings": 8192, "window_size": 4096, "dropout": 0.0, "intermediate_size": 14336, "bias": False, "rms_norm_eps": 1e-5}with torch.device("meta"): model = Mistral(cfg_7b)print(count(model)) # 7241732096Конфигурация
Заголовок раздела «Конфигурация»Пример из experiments/llm_only/configs/mistral_train.json:
| Параметр | Значение в примере | Смысл |
|---|---|---|
vocab_size | (из токенизатора) | размер словаря |
embed_dim | 256 | размерность эмбеддингов |
num_q_heads | 4 | число Query-голов |
num_kv_heads | 2 | число Key/Value-голов; num_q_heads должно делиться на него |
head_size | 64 | необязательный размер головы; по умолчанию embed_dim // num_q_heads (тогда embed_dim обязан делиться на num_q_heads). Если задан, num_q_heads · head_size может не совпадать с embed_dim; для RoPE — чётный |
num_layers | 4 | число блоков MistralDecoder |
max_position_embeddings | 512 | максимальная длина последовательности |
rms_norm_eps | (нет в примере) | необязательный eps всех RMSNorm, по умолчанию 1e-6; у Mistral 7B — 1e-5 |
rope_theta | (нет в примере) | необязательная база частот RoPE, по умолчанию 10000 — как в Mistral 7B v0.1; что она задаёт — в llama.md |
initializer_range | (нет в примере) | необязательное стандартное отклонение начальных весов Linear и Embedding, по умолчанию 0.02 — как в HF; см. training.md |
window_size | 16 | необязательная ширина скользящего окна внимания (окно — window_size + 1 позиций, см. выше); без ключа окна нет — обычное causal-внимание, как в Mistral 7B v0.2+ |
intermediate_size | (нет в примере) | необязательный скрытый размер SwiGLU, по умолчанию 4 · embed_dim; у Mistral 7B — 14336 (3.5·d) |
bias | (нет в примере) | необязательный: bias во всех Linear (Q/K/V, выход attention, три матрицы SwiGLU, голова), по умолчанию true; в Mistral 7B — false |
dropout | 0.1 | dropout после эмбеддингов и на выходах attention и FFN; в Mistral 7B dropout нет — для соответствия оригиналу 0 |
Все ключи используются конструктором Mistral.__init__. intermediate_size и bias меняют форму весов: по умолчанию сохранена прежняя структура, чтобы загружались старые чекпоинты. Неверные сочетания отклоняются с ValueError уже в конструкторе: embed_dim, не делящийся на num_q_heads без явного head_size, num_q_heads, не делящееся на num_kv_heads, нечётный head_size.
Отличия от Mistral 7B
Заголовок раздела «Отличия от Mistral 7B»Подробности, воспроизведение и варианты исправления — в бэклоге (номера пунктов в скобках).
| Mistral 7B | Здесь | |
|---|---|---|
| Скрытый слой SwiGLU | hidden_dim = 14336 при dim = 4096 (3.5·d) | 4·d по умолчанию; intermediate_size: 14336 — как в оригинале (30) |
| Bias | нет ни в одной проекции | во всех Linear по умолчанию; bias: false — как в оригинале (24) |
| Dropout | нет | после эмбеддингов, на выходах attention и SwiGLU (51); dropout: 0 убирает его полностью |
| Ширина окна | W + 1 позиций в тексте статьи и prefill эталона, W в HF | W + 1 (см. выше) |
eps RMSNorm | 1e-5 | 1e-6 по умолчанию, задаётся ключом rms_norm_eps |
| KV-кэш | кольцевой буфер (запись по pos % W) | torch.cat и обрезка срезом; результат тот же |
| Внимание в окне | ядра, считающие только полосу окна () | полная матрица и маска () |
| Префилл кусками | размером с окно | generate — весь промпт сразу; кусками — вручную через forward с кэшем |
Скользящее окно есть только в Mistral 7B v0.1 (sliding_window: 4096); в v0.2 и v0.3 его убрали (sliding_window: null в конфиге HF). Здесь это ключ window_size: без него окна нет.
Загрузка весов HuggingFace
Заголовок раздела «Загрузка весов HuggingFace»С ключами intermediate_size и "bias": false загружаются веса MistralForCausalLM — той же функцией convert_hf_state_dict, что у LLaMA (реэкспорт в llm.models.mistral); строки q_proj переставляются по num_attention_heads, k_proj — по num_key_value_heads.
from transformers import MistralForCausalLMfrom llm.models.mistral import Mistral, convert_hf_state_dict
hf = MistralForCausalLM.from_pretrained(...)c = hf.configconfig = {"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 or c.hidden_size // c.num_attention_heads, "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}if c.sliding_window is not None: config["window_size"] = c.sliding_window - 1 # окно здесь на позицию шире, см. «Ширина окна: W + 1»model = Mistral(config)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))Сверено со случайными MistralForCausalLM из transformers (со скользящим окном и без него, с head_dim, не равным hidden_size / num_attention_heads): логиты совпадают до ~1e-5, greedy-генерация с KV-кэшем дольше окна — токен в токен (llm/tests/models/test_mistral_mixtral_hf_parity.py). Настоящие веса (Mistral 7B — около 14 ГБ) для проверки слишком велики.
Mistral без window_size — это LLaMA с GQA, поэтому так же загружаются и чекпоинты LlamaForCausalLM с num_key_value_heads < num_attention_heads (Llama 2 70B и производные): проверено на случайной модели, логиты совпадают до ~1e-7.
Генерация
Заголовок раздела «Генерация»Mistral.generate(...) — унифицированная сигнатура BaseModel.generate (см. gpt.md и generation.md). Отличия от LLaMA — в кэше: при заданном window_size он не растёт дальше позиций на слой, а позиция следующего токена хранится в нём явно (next_pos). Позиции по-прежнему ограничены max_position_embeddings: когда текст длиннее, generate продолжает по последним токенам без кэша — как у всех моделей.
Что изменилось в Mixtral
Заголовок раздела «Что изменилось в Mixtral»- плотный
SwiGLU-FFN → Mixture-of-Experts: 8 параллельных SwiGLU-экспертов и роутер, на каждый токен работают 2 из них; - GQA, RoPE и RMSNorm остаются; блок декодера отличается только FFN-частью;
- скользящего окна в Mixtral 8x7B нет, а база RoPE увеличена до под контекст 32K (в репозитории
window_sizeу Mixtral остаётся необязательным ключом).
Подробности — в mixtral.md и mixture-of-experts.md.
Типичные ошибки и тонкости
Заголовок раздела «Типичные ошибки и тонкости»window_size = sliding_windowпри загрузке из HF. Окно окажется на позицию шире, логиты разойдутся; нужноsliding_window − 1.- Ожидание, что окно ускоряет обучение. Здесь маска накладывается на полную матрицу ; вычислений окно не экономит, экономит только кэш при генерации.
- Позиция по длине кэша. С окном длина кэша ≤ и не равна позиции; позицию берите из третьего элемента кэша (
next_pos), как делаетcache_start_pos. num_q_heads, не кратноеnum_kv_heads. Группы не получатся равными; конструктор бросаетValueError.- Путаница голов Q и K/V при перестановке строк.
q_projпереставляется по числу голов Q,k_proj— по числу голов K/V;convert_hf_state_dictбезnum_kv_headsдля GQA-чекпоинта даст неверные K. - Рецептивное поле ≠ контекст. может быть и больше, и меньше (у учебного конфига — 64 против 512).
- Mistral 7B = блок LLaMA + GQA (, ) + скользящее окно () + более широкий FFN (); по статье превосходит Llama 2 13B.
- GQA: головы Q делятся на групп с общими K/V; кэш и проекции K/V меньше в раз, вычисления внимания те же.
- Окно: маска ( позиций здесь, в HF); дальность через слои — , у Mistral 7B ≈ 131K.
- Кэш, ограниченный окном: в статье — кольцевой буфер с записью в , здесь —
torch.catи срез, кэш — тройка(K, V, next_pos). - Префилл кусками здесь работает через
forwardс кэшем для кусков любой длины. - для Mistral 7B; проверяется на
torch.device("meta").
Вопросы и упражнения
Заголовок раздела «Вопросы и упражнения»-
У модели , . Какую группу K/V читает голова Q с номером 5? А при ?
Ответ
, . При : , .
-
Сколько места займёт KV-кэш Mistral 7B во float16 для одного запроса длиной 16 384 токена без окна и с окном (по кольцевому буферу статьи)?
Ответ
128 КиБ на токен (см. Сколько экономит GQA). Без окна: КиБ ГиБ. С окном: КиБ МиБ, в 4 раза меньше.
-
Кольцевой буфер, , записаны позиции 0–10. В какой ячейке позиция 10 и какие позиции лежат в буфере?
Ответ
. В ячейках 0, 1, 2, 3 — позиции 8, 9, 10, 7.
-
Какое рецептивное поле у учебного конфига
mistral_train.json? Хватает ли его, чтобы последний токен при зависел от первого?Ответ
позиции. Не хватает: при последний токен зависит только от 64 предыдущих токенов входа (и от себя).
-
Сколько параметров сэкономил GQA в Mistral 7B по сравнению с MHA при тех же остальных размерах? Какая это доля модели?
Ответ
На слой: ; на 32 слоя — . MHA-вариант — параметров, экономия — 10 % от него.
-
В HF-конфиге
sliding_window = 4096. Какойwindow_sizeзадать здесь и сколько позиций тогда видит токен?Ответ
window_size = 4095; токен видит позиций — как в HF. -
Промпт из 12 токенов обрабатывается кусками по 4 в модели с . Какой формы матрица оценок внимания (без осей батча и голов) на каждом куске и сколько позиций в кэше после каждого?
Ответ
Кусок 1: кэша нет, оценки , после — кэш 4 позиции (0–3). Кусок 2: , кэш снова 4 (4–7). Кусок 3: , кэш — позиции 8–11,
next_pos = 12. -
Почему GQA не уменьшает число операций в , но всё равно ускоряет генерацию?
Ответ
Каждая из голов Q по-прежнему умножается на все ключи: на токен. Но при генерации по одному токену время упирается не в арифметику, а в чтение кэша из памяти, а кэш меньше в раз. Кроме того, освободившаяся память позволяет обрабатывать больший батч.
Литература
Заголовок раздела «Литература»Основная статья:
- Jiang et al. Mistral 7B. 2023. arXiv:2310.06825
Компоненты:
- Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. 2023. arXiv:2305.13245
- Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019. arXiv:1911.02150 — Multi-Query Attention
- Beltagy, Peters, Cohan. Longformer: The Long-Document Transformer. 2020. arXiv:2004.05150 — sliding window attention
- Touvron et al. LLaMA: Open and Efficient Foundation Language Models. 2023. arXiv:2302.13971 — базовая архитектура (RoPE, RMSNorm, SwiGLU)
- Touvron et al. Llama 2: Open Foundation and Fine-Tuned Chat Models. 2023. arXiv:2307.09288 — модели, с которыми сравнивается Mistral 7B