Запись архива

Jamba и гибриды SSM+attention: разбор архитектуры AI21 образца 2024 года

Разбираем paper Jamba от AI21 от 28 марта 2024 года: зачем смешивать Mamba, attention и MoE, что показали бенчмарки, где архитектура выигрывает и в чём её ограничения.

Схема архитектуры Jamba: один attention-слой на семь Mamba-слоёв, MoE в каждом втором слое и сравнение памяти KV-cache с Mixtral

28 марта 2024 года AI21 опубликовала paper Jamba: A Hybrid Transformer-Mamba Language Model и в тот же день выпустила официальный анонс модели. Этот разбор ограничен именно исторической версией Jamba-v0.1: открытой base-моделью 2024 года, а не более поздними Jamba 1.5, 1.6 или 2. Для практиков это важный кейс: Jamba пытается не заменить Transformer целиком, а уменьшить его слабые места на длинном контексте, добавив Mamba-слои и sparse-MoE.

Если коротко, AI21 показывает архитектурный компромисс: немного attention для точечной выборки по контексту, много SSM для дешёвой последовательной обработки и MoE для роста ёмкости без пропорционального роста активных параметров. Но почти все сильные выводы в статье основаны на внутренних измерениях авторов, поэтому ниже я отдельно отмечаю, где мы опираемся на факт из источника, а где начинается интерпретация.

Коротко

  • Jamba-v0.1 — это гибрид Transformer + Mamba + MoE с 12B активных и 52B общих параметров и контекстом 256K токенов, как указано в paper, model card и официальном анонсе.
  • В реализованной конфигурации paper описывает стек из четырёх Jamba-блоков по восемь слоёв с соотношением 1 attention на 7 Mamba; MoE включается в каждом втором слое с 16 экспертами и выбором top-2 на токен.
  • На общих бенчмарках Jamba по данным самой статьи близка к Mixtral 8x7B и местами к Llama-2 70B, но не выигрывает всё подряд: например, Mixtral выше на MMLU и BBH в таблице authors.
  • На long-context QA средний F1 в paper у Jamba составляет 0.44 против 0.43 у Mixtral. Прирост есть, но он небольшой и не выглядит как разгром.
  • Главный исследовательский вывод статьи не в «победе над Transformer», а в том, что гибрид в абляциях ведёт себя лучше, чем pure Mamba, на задачах, чувствительных к in-context learning и соблюдению формата ответа.

Контекст

Чтобы понять, зачем вообще понадобилась Jamba, нужно вспомнить исходные роли двух семейств моделей. Transformer сделал attention центральным механизмом работы с последовательностью: это дало сильное качество, но вместе с этим принесло дорогой inference на длинных контекстах и растущий KV-cache. В самом paper Jamba авторы прямо ставят проблему так: у attention-моделей длинный контекст быстро упирается и в память, и в пропускную способность.

Mamba, опубликованная 1 декабря 2023 года, предложила другой путь: selective state space model, где последовательность обрабатывается через состояние, а не через полный self-attention на каждом слое. Для длинных последовательностей это привлекательно, потому что вычисления и память масштабируются мягче. Но сами авторы Jamba пишут, что pure Mamba в их экспериментах уступала там, где модели нужно надёжно извлекать паттерн «вход → правильный формат выхода» и использовать few-shot-подсказки как рабочую память.

Отсюда и идея гибрида. Вопрос не «чем заменить Transformer», а «сколько attention достаточно, чтобы вернуть сильные стороны Transformer, не потеряв экономику long-context inference». Jamba — один из первых публично выпущенных ответов на этот вопрос в масштабе полноценной LLM, а не игрушечного абляционного прогона.

Метод

Как устроен блок Jamba

По paper, Jamba — это decoder-only архитектура, где внутри одного стека перемешаны attention-слои, Mamba-слои и sparse MoE. Базовая единица называется Jamba block. В опубликованной конфигурации paper описывает 4 блока по 8 слоёв каждый; внутри блока используется соотношение 1:7, то есть один attention-слой на семь Mamba-слоёв. Из этого следует, что в релизной конфигурации attention-слоёв всего четыре; это утверждение я беру именно из paper, а не из независимой верификации.

MoE в Jamba прикручивается не к attention, а к MLP-части слоя. В той же конфигурации MoE включается через слой (e=2), имеет 16 экспертов и роутер выбирает 2 эксперта на токен. Эта конструкция нужна для того, чтобы увеличить общую ёмкость модели, не делая каждый forward таким же дорогим, как dense-модель сопоставимого полного размера.

Практический итог конфигурации подтверждают сразу несколько официальных источников. Paper, model card и анонс AI21 сходятся в том, что опубликованная Jamba-v0.1 имеет 12B активных и 52B общих параметров, поддерживает 256K токенов контекста и рассчитана на размещение в одном 80GB GPU при 8-битной загрузке. Model card дополнительно уточняет прикладной предел: до 140K токенов на одном 80GB GPU в 8-bit. Это полезное различие: «модель поддерживает 256K» и «вам удобно держать 256K на одном устройстве в реальном пайплайне» — не одно и то же.

Есть и два менее очевидных инженерных решения из paper. Во-первых, authors пишут, что Jamba обходится без явного positional encoding и не использует RoPE в базовой конфигурации. Во-вторых, при масштабировании пришлось добавить внутренний RMSNorm в Mamba-слои, иначе на большом размере возникали spikes по loss.

Результаты

Ниже — сводка по таблицам 1–3 из paper Jamba. Это важно читать именно как результаты авторов на их конфигурации и их инфраструктуре, а не как независимый re-run. Для строки самой Jamba числа по нескольким общим бенчмаркам повторяются и в model card.

Модель Активные параметры KV-cache при 256K HellaSwag WinoGrande MMLU BBH Средний F1 на long-context QA
Mixtral 8x7B 12.9B 32GB 86.7 81.2 70.6 50.3 0.43
Jamba 12B 4GB 87.1 82.5 67.4 45.4 0.44

Что из этой таблицы стоит вынести на практике. Во-первых, Jamba действительно выглядит конкурентоспособной в своей размерной категории: по данным authors она немного выше Mixtral на HellaSwag, WinoGrande и среднем long-context F1. Во-вторых, это не история про тотальную победу. На MMLU и BBH Mixtral в той же статье выше. То есть hybrid SSM+attention здесь скорее меняет профиль компромисса, чем просто «делает всё лучше».

Long-context часть особенно показательная. В таблице 3 paper средний F1 у Jamba на пяти QA-наборах составляет 0.44 против 0.43 у Mixtral; Jamba выше на LongFQA, NarrativeQA и Natural Questions, но ниже на CUAD и SFiction. Это хороший пример того, почему long-context нельзя сводить к одной цифре окна контекста: даже при поддержке 256K выигрыш на реалистичных задачах оказался умеренным.

По эффективности авторы дают более агрессивные цифры. И paper, и официальный анонс утверждают, что на длинных контекстах Jamba даёт до 3x больше throughput, чем Mixtral 8x7B; paper уточняет сценарий: при контексте 128K токенов и генерации 512 выходных токенов на 4×A100. Те же источники также сходятся на утверждении, что Jamba помещает существенно больший контекст на одном A100 80GB, чем сопоставимые attention-only модели.

Самая интересная часть paper, на мой взгляд, — абляции. В таблице 6 authors показывают, что pure Mamba заметно проваливается на задачах, где нужно держать формат few-shot примеров и правильно продолжать шаблон ответа. По данным самой статьи на IMDB pure Attention даёт 84.1, pure Mamba — 48.8, а hybrid Attention-Mamba — 90.9; на NarrativeQA это 45.8, 27.7 и 43.7 соответственно. Авторы интерпретируют это как косвенное свидетельство того, что небольшой объём attention помогает восстановить in-context learning-поведение, которого pure SSM не хватает.

Интерпретация

Наш комментарий

Jamba важна не потому, что «доказала смерть Transformer». Наоборот: paper скорее показывает, что attention остаётся слишком полезным, чтобы выбрасывать его полностью. Архитектурный сдвиг здесь другой: attention можно сделать редким, а не повсеместным, если остальную работу по последовательности возьмут на себя SSM-слои.

Для инженеров это меняет прикладной вопрос. Вместо выбора между «чистым Transformer» и «чистым SSM» появляется третий путь: дозировать attention как дорогой, но стратегически важный механизм. В случае Jamba paper утверждает, что даже схема 1:7 уже возвращает достаточно свойств Transformer, чтобы гибрид не разваливался на задачах с few-shot-шаблоном.

Ещё один полезный вывод: спор об архитектуре нельзя отделять от режима эксплуатации. Если у вас RAG по длинным документам, агентные логи, юридические контракты или большие финансовые файлы, то важны не только MMLU и HellaSwag, но и то, как быстро и дёшево модель живёт на 64K, 128K и 256K контексте. В этом смысле Jamba — прежде всего статья об экономике inference, а уже потом об «альтернативе Transformer».

Ограничения

  • Почти все сильные выводы принадлежат самим авторам paper. В статье нет независимого воспроизведения throughput и long-context QA. Поэтому цифры вроде 3x throughput или преимущества на среднем F1 стоит читать как авторское измерение на конкретном стеке, а не как универсальную константу.
  • Релизная Jamba-v0.1 — это base model, а не готовый чат-продукт. Paper прямо предупреждает, что модель не проходила alignment или instruction tuning и не должна использоваться с конечными пользователями без дополнительной адаптации. Анонс AI21 повторяет ту же мысль: guardrails и дообучение остаются на стороне разработчика.
  • Прирост на long-context QA небольшой. Разница 0.44 против 0.43 в среднем F1 выглядит положительно, но не меняет класс модели. Кроме того, Jamba не выигрывает все пять наборов: на CUAD и SFiction в той же таблице она ниже Mixtral.
  • Синтетический long-context тест не равен реальному reasoning на 256K. Needle-in-a-haystack показывает, что модель может извлечь спрятанную строку из длинного контекста. Это полезно, но не доказывает устойчивое многошаговое рассуждение, надёжное цитирование и отсутствие деградации на произвольных задачах при 256K.
  • Есть важные аппаратные оговорки. Model card пишет, что для нормальной работы нужны CUDA, mamba-ssm, causal-conv1d и современный transformers; при BF16/FP16 модель уже не влезает в один 80GB GPU, а практический single-GPU предел 140K относится именно к 8-bit загрузке. Для прод-инференса это не мелочь, а часть архитектурной цены.
  • Историческая рамка важна. После марта 2024 года AI21 выпустила следующие поколения семейства Jamba, но этот paper не доказывает их свойства автоматически. Переносить выводы из Jamba-v0.1 на поздние версии без отдельной проверки было бы ошибкой.

Вывод

Jamba 2024 года — это не финальный ответ на вопрос «что придёт после Transformer», а аккуратное доказательство того, что гибрид SSM+attention вообще может работать в масштабе открытой LLM. Сильная сторона работы — не рекорд по любому одному бенчмарку, а демонстрация полезного компромисса между качеством, памятью и скоростью на длинном контексте.

Если вы практик, главный урок из paper такой: смотреть нужно не на ярлык «SSM» или «Transformer», а на конфигурацию компромисса. В Jamba эту роль играют редкий attention, Mamba как основа последовательной обработки и MoE как способ нарастить ёмкость без dense-цены на каждый токен.

Источники

FAQ

Что такое Jamba в одном предложении?

Это LLM AI21, где редкие attention-слои, частые Mamba-слои и sparse-MoE объединены в один decoder для более дешёвого long-context inference.

Почему AI21 не отказалась от attention полностью?

По логике authors, pure SSM хорошо помогает с масштабированием по длине последовательности, но хуже справляется с in-context learning и точечным извлечением шаблона ответа. Небольшое число attention-слоёв должно компенсировать именно это.

Можно ли считать 256K доказательством сильного long-context reasoning?

Нет. 256K означает поддерживаемый размер контекста, но не гарантирует одинаково хорошее качество на всех задачах у верхней границы окна. В самой статье реалистичный long-context прирост над Mixtral есть, но он умеренный.

Что здесь важнее для инженера: SSM, MoE или отношение 1:7?

Важна вся связка. SSM влияет на память и throughput, attention — на recall и few-shot поведение, MoE — на ёмкость модели, а соотношение 1:7 задаёт баланс между ними.

Читайте также