Jamba: A Hybrid Transformer-Mamba Language Model

Opher Lieber, Barak Lenz, Hofit Bata, Gal Cohen, Jhonathan Osin, Itay Dalmedigos, Erez Safahi, Shaked Meirom, Yonatan Belinkov, Shai Shalev-Shwartz, Omri Abend, Raz Alon, Tomer Asida, Amir Bergman, Roman Glozman, Michael Gokhman, Avashalom Manevich, Nir Ratner, Noam Rozen, Erez Shwartz, Mor Zusman, Yoav Shoham

Introduction

We introduce Jamba, a new publicly available large language model. Jamba is based on a novel hybrid architecture, which combines Transformer layers with Mamba layers , a recent state-space model , as well as a mixture-of-experts (MoE) component . Jamba thus combines two orthogonal architectural designs that together give it improved performance and higher throughput, while maintaining a manageable memory footprint. The 7B-based Jamba model (12B active parameters, 52B total available parameters) we are releasing was designed to fit in a single 80GB GPU, but the Jamba architecture supports other design choices, depending on one’s hardware and performance requirements.

The fundamental novelty of Jamba is its hybrid Transformer-Mamba architecture (though see mention below of recent related efforts). Despite the immense popularity of the Transformer as the predominant architecture for language models, it suffers from two main drawbacks. First, its high memory and compute requirements hinders the processing of long contexts, where the key-value (KV) cache size becomes a limiting factor. Second, its lack of a single summary state entails slow inference and low throughput, since each generated token performs a computation on the entire context. In contrast, older recurrent neural network (RNN) models, which summarize an arbitrarily long context in a single hidden state, do not suffer from these limitations. RNN models have their own shortcomings, however. They are costly to train since training cannot be parallelized across time steps. And they struggle with long distance relationships, which the hidden state captures to only a limited extent.

Recent state space models (SSMs) like Mamba are more efficient to train than RNNs and are more capable at handling long distance relationships, but still lag behind the performance of comparably sized Transformer language models. Taking advantage of both model families, Jamba combines Transformer and Mamba layers, at a certain ratio. Varying the ratio of Transformer/Mamba layers allows balancing memory usage, efficient training, and long context capabilities.

A few other recent attempts to combine Attention and SSM modules are worth noting. mixes an S4 layer with a local attention layer, followed by a sequence of local attention layers; it shows experiments with small models and simple tasks. reports that interleaving Mamba and attention layers is only slightly better than pure Mamba in terms of perplexity, with models up to 1.3B parameters. starts with an SSM layer followed by chunk-based Transformers, with models up to 1.3B showing improved perplexity. adds an SSM layer before the self-attention in a Transformer layer, while adds the SSM after the self-attention, both showing improvements on speech recognition. replaces the MLP layers in the Transformer by Mamba layers, and shows benefits in simple tasks. These efforts are different from Jamba both in the particular way in which the SSM component is mixed with the attention one, and in the scale of implementation. Closest are perhaps H3 , a specially designed SSM that enables induction capabilities, and a generalization called Hyena . The former proposed a hybrid architecture that replaces the second and middle layers with self-attention, and was implemented with up to 2.7B parameters and 400B training tokens. However, as shown in , its perfomance lags that of pure Mamba. Based on Hyena, StripedHyena interleaves attention and SSM layers in a 7B parameter model. However, it lags behind the Attention-only Mistral-7B . All of this renders Jamba the first production-grade Attention-SSM hybrid model. Scaling the hybrid Jamba architecture required overcoming several obstacles, which we dicsuss in Section 6.

Jamba also includes MoE layers , which allow increasing the model capacity (total number of available parameters) without increasing compute requirements (number of active parameters). MoE is a flexible approach that enables training extremely large models with strong performance . In Jamba, MoE is applied to some of the MLP layers. The more MoE layers, and the more experts in each MoE layer, the larger the total number of model parameters. In contrast, the more experts we use at each forward pass, the larger the number of active parameters as well as the compute requirement. In our implementation of Jamba, we apply MoE at every other layer, with 16 experts and the top-2 experts used at each token (a more detailed discussion of the model architecture is provided below).

We evaluated our implementation of Jamba on a wide range of benchmarks and found it performs comparably to Mixtral-8x7B , which has a similar number of parameters, and also to the larger Llama-2 70B . In addition, our model supports a context length of 256K tokens – the longest supported context length for production-grade publicly available models. On long-context evaluations, Jamba outperformes Mixtral on most of the evaluated datasets. At the same time, Jamba is extremely efficient; for example, its throughput is 3x that of Mixtral-8x7B for long contexts. Moreover, our model fits in a single GPU (with 8bit weights) even with contexts of over 128K tokens, which is impossible with similar-size attention-only models such as Mixtral-8x7B.

Somewhat unusually for a new architecture, we release Jamba (12B active parameters, 52B total available parameters) under Apache 2.0 license: https://huggingface.co/ai21labs/Jamba-v0.1. We do so since we feel that the novel architecture of Jamba calls for further study, experimentation, and optimization by the community. Our design was based on various ablation experiments we conducted to explore the effect of different tradeoffs and design choices, and insights gleaned from those. These ablations were performed at scales of up to 7B parameters, and training runs of up to 250B tokens. We plan to release model checkpoints from these runs.

Important notice: The Jamba model released is a pretrained base model, which did not go through alignment or instruction tuning, and does not have moderation mechanisms. It should not be used in production environments or with end users without additional adaptation.

Model Architecture

Jamba is a hybrid decoder architecture that mixes Transformer layers with Mamba layers , a recent state-space model (SSM) , in addition to a mixture-of-experts (MoE) module . We call the combination of these three elements a Jamba block. See Figure 1 for an illustration.

Combining Transformer, Mamba, and MoE elements allows flexibility in balancing among the sometimes conflicting objectives of low memory usage, high throughput, and high quality. In terms of memory usage, note that comparing the total size of the model parameters can be misleading. In an MoE model, the number of active parameters that participate in any given forward step may be much smaller than the total number of parameters. Another important consideration is the KV cache – the memory required to store the attention keys and values in the context. When scaling Transformer models to long contexts, the KV cache becomes a limiting factor. Trading off attention layers for Mamba layers reduces the total size of the KV cache. Our architecture aims to provide not only a small number of active parameters but also an 8x smaller KV cache compared to a vanilla Transformer. Table 1 compares Jamba with recent publicly available models, showing its advantage in maintaining a small KV cache even with 256K token contexts.

In terms of throughput, with short sequences, attention operations take up a small fraction of the inference and training FLOPS . However, with long sequences, attention hogs most of the compute. In contrast, Mamba layers are more compute-efficient. Thus, increasing the ratio of Mamba layers improves throughput especially for long sequences.

Here is a description of the main configuration, which provides improved performance and efficiency. Section 6 contains results from ablation experiments supporting the design choices.

The basic component is a Jamba block, which may be repeated in sequence. Each Jamba block is a combination of Mamba or Attention layers. Each such layer contains either an attention or a Mamba module, followed by a multi-layer perceptron (MLP). The different possible types of layers are shown in Figure 1(b).The figure shows a potential Attention MoE layer, which our architecture does not use, but future variants could. A Jamba block contains ll layers, which are mixed at a ratio of a:ma:m, meaning aa attention layers for every mm Mamba layers.

In Jamba, some of the MLPs may be replaced by MoE layers, which helps increase the model capacity while keeping the active number of parameters, and thus the compute, small. The MoE module may be applied to MLPs every ee layers. When using MoE, there are nn possible experts per layer, with a router choosing the top KK experts at each token. In summary, the different degrees of freedom in the Jamba architecture are:

a:ma:m: ratio of attention-to-Mamba layers.

ee: how often to use MoE instead of a single MLP.

KK: number of top experts used at each token.

Given this design space, Jamba provides flexibility in preferring certain properties over others. For example, increasing mm and decreasing aa, that is, increasing the ratio of Mamba layers at the expense of attention layers, reduces the required memory for storing the key-value cache. This reduces the overall memory footprint, which is especially important for processing long sequences. Increasing the ratio of Mamba layers also improves throughput, especially at long sequences. However, decreasing aa might lower the model’s capabilities.

Additionally, balancing nn, KK, and ee affects the relationship between active parameters and total available parameters. A larger nn increases the model capacity at the expense of memory footprint, while a larger KK increases the active parameter usage and the compute requirement. In contrast, a larger ee decreases the model capacity, while decreasing both compute (when KK>1) and memory requirements, and allowing for less communication dependencies (decreasing memory transfers as well as inter-GPU communication during expert-parallel training and inference).

Jamba’s implementation of Mamba layers incorporate several normalizations that help stabilize training in large model scales. In particular, we apply RMSNorm in the Mamba layers.

We found that with the Mamba layer, positional embeddings or mechanisms like RoPE are not necessary, and so we do not use any explicit positional information.

Other architecture details are standard, including grouped-query attention (GQA), SwiGLU activation function , and load balancing for the MoE . The vocabulary size is 64K. The tokenizer is trained with BPE and each digit is a separate token . We also remove the dummy space used in Llama and Mistral tokenizers for more consistent and reversible tokenization.

Reaping the Benefits

The specific configuration in our implementation was chosen to fit in a single 80GB GPU, while achieving best performance in the sense of quality and throughput. In our implementation we have a sequence of 4 Jamba blocks. Each Jamba block has the following configuration:

a:m=1:7a:m=1:7: ratio attention-to-Mamba layers.

e=2e=2: how often to use MoE instead of a single MLP.

K=2K=2: number of top experts used at each token.

The a:m=1:7a:m=1:7 ratio was chosen according to preliminary ablations, as shown in Section 6, since this ratio was the most compute-efficient variant amongst the best performing variants in terms of quality.

The configuration of the experts was chosen to enable the model to fit in a single 80GB GPU (with int8 weights), while including sufficient memory for the inputs. In particular, nn and ee were balanced to have an average of ∼\sim8 experts per layer. In addition, we balanced nn, KK, and ee to allow for high quality, while keeping both compute requirements and communication dependencies (memory transfers) checked. Accordingly, we chose to replace the MLP module with MoE on every other layer, as well as have a total of 16 experts, two of which are used at each token. These choices were inspired by prior work on MoE and verified in preliminary experiments.

Figure 2 shows the maximal context length that fits a single 80GB GPU with our Jamba implementation compared to Mixtral 8x7B and Llama-2-70B. Jamba provides 2x the context length of Mixtral and 7x that of Llama-2-70B.

Overall, our Jamba implementation was successfully trained on context lengths of up to 1M tokens. The released model supports lengths of up to 256K tokens.

2 Throughput Analysis

For concreteness, we present results of the throughput in two specific settings.Referring to end-to-end throughput (encoding+decoding). The results should be taken relatively rather than absolutely, as they are without possible optimizations. In the first setting, we have varying batch size, a single A100 80 GB GPU, int8 quantization, 8K context length, generating output of 512 tokens. As Figure 3(a) shows, Jamba allows processing of large batches, leading to a 3x increase in throughput (tokens/second) over Mixtral, which does not fit with a batch of 16 despite having a similar number of active parameters.

In the second setting, we have a single batch, 4 A100 GPUs, no quantization, varying context lengths, generating output of 512 tokens. As demonstrated in Figure 3(b), at small context lengths all models have a similar throughput. Jamba excels at long contexts; with 128K tokens its throughput is 3x that of Mixtral. Note that this is despite the fact that Jamba has not yet enjoyed optimizations of the kind the community has developed for pure Transformer models over the past six years. We can expect the throughut gap to increase as such optimizations are developed also for Jamba.

Training Infrastructure and Dataset

The model was trained on NVIDIA H100 GPUs. We used an in-house proprietary framework allowing efficient large-scale training including FSDP, tensor parallelism, sequence parallelism, and expert parallelism.

Jamba is trained on an in-house dataset that contains text data from the Web, books, and code, with the last update in March 2024. Our data processing pipeline includes quality filters and deduplication.

Evaluation

In general we approach benchmarks cautiously, as they correlate only partially with what matters in real applications, and furthermore invite gaming the system in order to boast vanity numbers. Nevertheless, we present several indicative results.

We report results with a wide range of standard academic benchmarks:

HellaSwag (10-shot) , WinoGrande (5-shot) , ARC-E (0-shot) and ARC-Challenge (25-shot) , and PIQA (zero-shot) .

GSM8K (3-shot CoT) , HumanEval (pass@1) , Natural Questions closed-book (NQ; 5-shot) , and TruthfulQA (zero-shot) .

Table 2 compares Jamba to several publicly available models on common academic benchmarks for evaluating language models. We compare with Llama-2 13B , which has about the same number of active paramters as our model, Llama-2 70B, which is larger than our model, Gemma , which has 7B parameters, and Mixtral , which has about the same number of active and total parameters as our model.

Noticebly, Jamba performs comparably to the leading publicly available models of similar or larger size, including Llama-2 70B and Mixtral. At the same time, our model has a smaller number of total available parameters than Llama-2 (52B compared to 70B). Moreover, as a sparse model, Jamba has only 12B active parameters, similar to Mixtral’s 12.9B active parameters. However, as a fully-attentional model, Mixtral has a large memory footprint with long sequences, requiring 32GB for the KV cache with 256K tokens. In contrast, thanks to its hybrid Attention-Mamba architecture, Jamba’s KV cache takes only 4GB even at such a long context (Section 2). Importantly, our Jamba achieves such a strong performance while having much better throughput than Llama-2 70B and Mixtral, up to 3x improvement (Section 3.2).

In summary, Jamba demostrates the ability of hybrid architectures to reach the performance of state-of-the-art Transformer based models of the same size class, while having the benefits of an SSM.

2 Long-Context Evaluations

We have successfully trained Jamba models with context lengths of up to 1M tokens. The released model handles context lengths of up to 256K tokens. In this section, we evaluate it on synthetic and naturalistic benchmarks that test is long-context capabilities.

As Figure 4 shows, Jamba has excellent performance in the needle-in-a-haystack evaluation, which requires retrieving a simple statement planted in a long context window . This result is noteworthy especially given that our implementation of Jamba uses only 4 attention layers.

2.2 Naturalistic long-context evaluation

We evaluate Jamba’s ability to handle long contexts using question-answering benchmarks, consisting of long inputs. To this end, we repurpose five of the longest-context datasets from L-Eval , by structuring them in a few-shot format (we use 3-shots in all experiments here). Specifically, we evaluated the models on the following datasets: NarrativeQA (QA on narratives; ), LongFQA (finance; ), Natural Questions (NQ; Wikipedia; ), CUAD (law; ), and SFiction (science fiction). The average input length in these datasets ranges from 6K to 62K tokens. These context lengths are further highly expanded by the few-shot format.

Table 3 summarizes the evaluation results, in terms of F1. Jamba outperforms Mixtral on most of the datasets as well as on average. In addition, as these long-context tasks require substantial computation, here Jamba’s efficiency shines, with much better throughput with long contexts (Section 3.2).

Ablations and Insights

This section discusses ablation experiments we ran for different design choices in our implementation of the Jamba architecture. First we show the benefit of combining attention and Mamba layers, at which ratio they should be combined, and how to interleave them. We investigate cases where pure Mamba fails, suggesting that it struggles to develop in-context learning capabilities, while the Attention–Mamba hybrid exhibits in-context learning similar to vanilla Transformers. Then we show the benefit of adding MoE on top of a hybrid Attention–Mamba model. Finally, we share two additional learnings that we found useful: explicit positional information is not needed in Jamba, and Mamba layers necessitate special normalization to stabilize training at large scale.In all the ablation experiments, “pure Mamba” refers to models with Mamba layers interleaved with MLP layers.

For these ablations, we report the following measures, which exhibit informative performance even at small data or model scale.

Academic benchmarks: HellaSwag (10-shot) , WinoGrande (5-shot) , Natural Questions closed-book (NQ; 5-shot) .

HuggingFace OpenLLM leaderboard (OLLM) : a summary statistic of several datasets. We report results with our reproduction.

Perplexity evaluations: we report log-prob (per byte) on texts from three domains: C4, Books, and code.

We first investigate the ratio of Attention to Mamba layers (a:ma:m), with 1.3B parameters models trained for 250B tokens. As Table 4 shows, the hybrid Jamba model outperforms the pure attention or Mamba models. The ratio of attention-to-Mamba layers may be 1:3 or 1:7 with virtually no performance difference. Figure 5 shows the training loss of these models, where Jamba exhibits improved loss during training. Given that a 1:7 ratio is more compute-efficient and shows similar performance, we opt for it in our larger-scale experiments.

Next, we compare performance of vanilla Transformer, vanilla Mamba, and Attention-Mamba hybrid models, at 7B model size, after training on 50B tokens. As Table 5 shows, the pure Mamba layer is quite competitive, but lags slightly behind pure Attention. The hybrid Attention-Mamba (without MoE) outperforms the pure models while obtaining better throughput than the vanilla Transformer (Section 3.2).

Figure 6 shows the training loss of the three architectures. While the pure Transformer and Mamba models have a similar convergence, the hybrid Jamba (no MoE) has a lower loss throughout this run.

2 Why does the Combination Work?

The pure Mamba model showed fairly good results in most tasks early on, including in general perplexity evaluations. However, it performed substantially worse than the pure Attention model in three common benchmark tasks: IMDB , QuAC , and NarrativeQA . In contrast, the hybrid Attention-Mamba performed similarly to the Attention model on these datasets. Table 6 shows the results for 1.3B models after 250B tokens.

Looking into these results further, we found out that the pure Mamba model often does not follow the correct format. For instance, in the IMDB dataset, answer choices are “Positive” or “Negative”. While the Attention model adheres to this format, the pure Mamba model often produces other answers, such as “Very Good”, “Very Positive”, “Funny”, “Bad”, “Poor”, and “3/10”. While these may be considered correct answers, the difficulty of Mamba to adhere to the format suggests a potential problem. Indeed, to perform successful in-context learning, it is important for models to capture the input-output format . The hybrid Attention-Mamba model follows the format successfully, just like the pure Attention model.

We hypothesize that this phenomenon points to a limitation of SSMs – a potential difficulty in in-context learning (ICL). Indeed, the ability to perform ICL has been linked to the emergence of so-called induction heads in Transformer language models during training, which perform approximate copying operations that are supportive of ICL . We conjecture that the lack of an attention mechanism in the pure Mamba model makes it difficult for it to learn in-context. While Mamba may learn to copy and perform simple ICL when explicitly trained to do so (, it is not clear if ICL is an emergent capability in SSM as is typical of Transformer models. In contrast, the hybrid Attention–Mamba model does perform successful ICL, even when only 1 out of 8 layers is an Attention one.

As anecdotal evidence of an emergent induction mechanism, we visualize in Figure 7 the attention of an example head from a 1.3B Attention-Mamba hybrid model (no MoE), on an IMDB example where the pure Mamba failed and the hybrid succeeded. Clearly, the attention from the last token (“:”) is focused on the labels from the few-shot examples. We have found 12 such heads in our hybrid model, in all three attention layers (which correspond to layers 4, 12, 20 in the model).

Future work can further investigate the emergence of ICL in hybrid models at large scale. Our released checkpoints would hopefully facilitate such investigations. Finally, very recent work has attempted to extract attention-like scores from state-space models like Mamba , which opens another direction to search for induction capabilities in state-space models.

3 The Effect of Mixture-of-Experts (MoE)

Recent work has shown that MoE improves Transformer language models while keeping compute manageable .There is also initial evidence that MoE helps Mamba layers, albeit at small model and data scale . However, it is not clear if MoE integrates well with state-space models at a large scale, and specifically with our hybrid Attention–Mamba architecture. Indeed, Table 7 shows that MoE improves the performance of the hybrid Attention-Mamba architecture at large scale (7B parameters trained on 50B tokens). The MoE variant has n=16n=16 total experts, K=2K=2 experts used at each token, and MoE is applied every e=2e=2 layers, as described in Section 3.1.

4 Stabilizing Mamba at large scale

When training Jamba models of up to 1.3B parameters, we observed stable training without special problems. However, when scaling to the largest model released here (7B-based, which has 12B/52B active/total parameters), we encountered large loss spikes. Investigating this revealed that inner parts of the Mamba layers suffer from large activation values, leading to the spikes. We therefore added RMSNorm to internal activations. As Figure 8 shows, this stabilized training and prevented additional loss spikes.

5 Jamba does not Require Explicit Positional Information

Table 8 shows results of the Jamba architecture (with MoE) with no positional information and when applying RoPE in the attention layers (1.3B parameter models, 250B tokens). The results are similar, suggesting that explicit positional information may not be required for the hybrid architecture. Presumably, the Mamba layers, which are placed before attention layers, provide implicit position information.Some prior evidence suggested that Transformer decoder models do not need positional encodings . However, all existing large scale models do use some sort of explicit position information.

Conclusion

We presented Jamba, a novel architecture which combines Attention and Mamba layers, with MoE modules, and an open implementation of it, reaching state-of-the-art performance and supporting long contexts. We showed how Jamba provides flexibility for balancing performance and memory requirements, while maintaining a high throughput. We experimented with several design choices such as the ratio of Attention-to-Mamba layers and discussed some discoveries made during the development process, which will inform future work on hybrid attention–state-space models. To facilitate such research, we plan to release model checkpoints from smaller-scale training runs. The largest model we provide with this release has 12B active and 52B total available parameters, supporting context lengths of up to 256K tokens and fitting in a single 80GB GPU even when processing 140K-token texts.

References