Jamba-1.5: Hybrid Transformer-Mamba Models at Scale
Jamba Team, Barak Lenz, Alan Arazi, Amir Bergman, Avshalom Manevich, Barak Peleg, Ben Aviram, Chen Almagor, Clara Fridman, Dan Padnos, Daniel Gissin, Daniel Jannai, Dor Muhlgay, Dor Zimberg, Edden M Gerber, Elad Dolev, Eran Krakovsky, Erez Safahi, Erez Schwartz, Gal Cohen, Gal Shachaf, Haim Rozenblum, Hofit Bata, Ido Blass, Inbal Magar, Itay Dalmedigos, Jhonathan Osin, Julie Fadlon, Maria Rozman, Matan Danos, Michael Gokhman, Mor Zusman, Naama Gidron, Nir Ratner, Noam Gat, Noam Rozen, Oded Fried, Ohad Leshno, Omer Antverg, Omri Abend, Opher Lieber, Or Dagan, Orit Cohavi, Raz Alon, Ro'i Belson, Roi Cohen, Rom Gilad, Roman Glozman, Shahar Lev, Shaked Meirom, Tal Delbari, Tal Ness, Tomer Asida, Tom Ben Gal, Tom Braude, Uriya Pumerantz, Yehoshua Cohen, Yonatan Belinkov, Yuval Globerson, Yuval Peleg Levy, Yoav Shoham
Introduction
This paper introduces Jamba-1.5, two new large language models based on our Jamba architecture , which are available for public use. Jamba-1.5-Mini is an updated and instruction-tuned version of our earlier Jamba release . Like its smaller sibling, Jamba-1.5-Large is a hybrid architecture that mixes Transformer and Mamba layers, with a mixture-of-experts (MoE) module . Since the introduction of Jamba, similar efforts have confirmed the benefits of combining Transformer and state-space-models at a scale of up to 8B parameters . Jamba-1.5-Large demonstrates the benefits of this architecture at a much larger scale. It has 94B active parameters, out of a total of 398B parameters. Even at this large size, the model can fit on a single machine with 8 80GB GPUs when processing a context of 256K tokens, thanks to the efficiency of the Jamba architecture in addition to a novel quantization technique we have developed, as described in Section 3.1.
Both Jamba-1.5-Mini and Jamba-1.5-Large are instruction-tuned models, having undergone post-training to provide them with various capabilities. Our evaluations across a wide range of benchmarks show that they perform comparably to models at their size, while offering the efficiency benefits of the Jamba architecture. In particular, Jamba-1.5 models shine at long-context evaluations, making them the only models with an effective length of 256K on the RULER benchmark, while offering 10x reduction in KV cache memory as well as superior throughput and latency.
We make the Jamba-1.5 models available under the Jamba Open Model License: https://www.ai21.com/licenses/jamba-open-model-license. The models are publicly available: Jamba-1.5-Mini: https://huggingface.co/ai21labs/AI21-Jamba-1.5-Mini Jamba-1.5-Large: https://huggingface.co/ai21labs/AI21-Jamba-1.5-Large
Model Architecture
Jamba-1.5-Large is based on Jamba , our hybrid decoder architecture that mixes Transformer layers with Mamba layers , a state-space model (SSM) , in addition to a mixture-of-experts (MoE) module . See for a detailed description of this architecture.
During our work on Jamba , we found that the combination of Transformer, Mamba, and MoE elements facilitates balancing desiderata of throughput, memory usage, and quality. Jamba-1.5-Large demonstrates this flexibility at a larger scale.
Jamba-1.5-Large follows the same Jamba structure but with a larger capacity. It has 94B active parameters and 398B total parameters. It has blocks, with each block having the following specs:
ratio of attention-to-Mamba layers. This ratio was found optimal in our work on Jamba and similar ratios was also confirmed as successful in follow-up work .
MoE is used instead of a single MLP every layers. There are experts, and we select the top at each token.
The hidden state dimensionality is .
The number of attention query heads is and the number of KV heads is .
Table 1 compares the Jamba-1.5 models to publicly available models of similar sizes. Jamba-1.5-Mini has a similar number of active parameters as Mixtral 8x7B, while Jamba-1.5-Large’s active parameter count is between LLaMA-3.1-70B and Mistral-Large-2. At the same time, both our Jamba models have a much smaller KV cache memory usage (at 256K tokens) compared to all other models, with roughly an order of magnitude reduction compared to their respective counterparts.
With these settings, and our specialized quantization (Section 3.1), Jamba-1.5-Large can be served on a single machine with 8 80GB GPUs with context lengths up to 256K tokens.
For this release, we experimented also with Mamba-2 , a faster and improved version of Mamba, which was reported to outperform Mamba and Transformers separately. However, as Figure 1 shows, we found that in a hybrid architecture, the Mamba-1-Attention combination works better than Mamba-2-Attention, so we use Mamba-1 in Jamba-1.5-Large. (We also found the hybrid architecture to outperform pure Mamba-2.) We hypothesize this is because some of the advantages of Mamba-2 over Mamba-1, in particular the ability to use a much larger state size, are less significant when we have full attention layers interleaved between the Mamba layers, as they can pool information from the entire context.
Serving Considerations and Improvements
We share a few insights and improvements we have introduced to allow for efficient serving of Jamba models at a large scale.
To support efficient serving of Jamba-1.5-Large, we developed a new quantization technique, which we dub ExpertsInt8. We observe that over 85% of the model weights are in the MoE layers, and over 90% are in MoE or MLP layers. We wish to quantize these weights while still enjoying the benefits of fast BF16 kernels. To do so, we quantize the MoE and MLP weights to INT8, save them in INT8, and dequnatize them back to BF16 before the actual computation. Importantly, the dequantization step happens directly inside the fused_moe kernel in vLLM . In this way, the dequantization process adds negligible overhead, and even leads to improved latency over BF16.We attribute this to the the kernel operating on relatively small blocks of weights and activations, which it moves from GPU HBM to SRAM prior to performing the computations. In our implementation, the weights move from HBM to SRAM when they are in int8, so it takes less time as their memory footprint is cut by half. We have contributed our modified fused_moe kernel to vLLM.Pull request here: https://github.com/vllm-project/vllm/pull/7415
Our ExpertsInt8 method has several advantages. First, it is fast; quantization only takes a few seconds at model loading. Second, unlike most other techniques in vLLM, it does not rely on calibration, which can take hours or days and can be unstable. Third, we can still use BF16 to hold large activations. Fourth, it is available to use on A100 GPUs, unlike FP8, which is only available on H100. Finally, our quantization matches FP8 in latency, while surpassing other quantization techniques, without a loss in quality.
Figure 2 compares the latency with different quantization techniques using Jamba-1.5-Mini, Jamba-1.5-Large, and two Mixtral models (8x78B and 8x22B). On H100 GPUs, ExpertsInt8 matches the latency of FP8. On A100, where FP8 is unavailable, ExpertsInt8 is an attractive technique, outperforming GPTQ by a large margin. Together with the advtanages of ExpertsInt8 explained above, this makes it an attractive quantization technique for serving large MoE models.
2 Activation Loss
During pre-training, we found that certain activations, namely outputs of specific experts as well as the the output of the last Mamba layers, were gradually increasing in magnitude for certain input tokens, eventually reaching values as high as . Although we did not find this to hurt the pre-training itself, which was done in BF16 precision, the magnitude of the activations could cause numerical issues during inference as some quantization libraries support only FP16 precision for activations, which has a maximum range of 64K.
To alleviate these concerns, we added an “Activation Loss” term, proportional to the mean-square of activations in the forward pass, with a configurable factor, which penalizes larger activation values. We found via experimentation that this auxilary loss has no affect on the training even with values up to at least . For Jamba-1.5-Large, we used which was enough to reduce the activations to an acceptable range (2K-3K max). Moreover, adding this auxilary loss reduced the activations almost instantly, allowing it to be added only towards the end of the training without any affect on training speed and quality.
To validate this approach, we ran our full evaluation suite on the model using FP16 activations and obtained the same results as the BF16 evaluations without any nans/overflows.
Throughput and Latency Analysis
Thanks to the hybrid Jamba architecture, our Jamba-1.5 models provide excellent throughput and latency. Figures 3 and 4 show this for Jamba-1.5-Mini and Jamba-1.5-Large, respectively. As shown in the figures, our models obtain much better latency and throughput than similarly-sized models. Their advantage shines at long contexts, with substantial gaps. Importantly, Jamba-1.5-Large runs efficiently even at long contexts, where the large LLaMA3-405B cannot run on the same hardware.
Training
Jamba-1.5-Large was trained on NVIDIA H100 GPUs using our in-house proprietary framework, which includes FSDP, tensor parallelism, sequence parallelism, and expert parallelism. For the latter we have adapted MegaBlocks .
2 Training Stages
The model was trained in three stages. During pre-training, it was first trained on an in-house dataset last updated in March 2024. Our pre-training dataset is a mixture of publicly available web documents, code, books and scientific articles. Our pre-processing pipeline includes parsing, quality filters, and deduplication. To make the best use of publicly available data, we developed our own in-house parser, and used it to extract text and formatting. The exact data mixture was determined through various ablations. This stage included multilingual data with emphasis on the following languages: English, Spanish, French, Portueguse, Italian, Dutch, German, Arabic, and Hebrew. It was then trained for a short phase of mid-training with a high proportion of long documents to emphasize its long-range capabilities. Finally, the model went through post-training, described in the next section.
3 Post-training
Our approach to post-training aims to achieve two objectives simultaneously: (i) provide the model with various skills and conversational capabilities; (ii) retain capabilities from pre-training and especially the long-context capabilities from mid-training. These two objectives are partly conflicting, since most of the available post-training datasets consist of relatively short examples.
Given these considerations, our post-training process involves supervised fine-tuning on high-quality conversational data, skill-specific data, and long-context data. Mixing these different types of data aims to retain long-context capabilities and acquire desired skills. As shown in the evaluations below, we find that our models perform very well in long-context evaluations.
When performing supervised fine-tuning, we make heavy use of synthetic data, as is common in recent foundation models and reflecting our approach for constructing structured data for building compound AI systems . We developed multiple different data synthesis pipelines, targeting different model capabilities. All pipelines apply the following pattern: (i) Sample or generate prompts in a target distribution; (ii) Generate responses from language models; (iii) Filter or rank responses by quality according to automatic validation and scoring; and (iv) Post-edit to remove artifacts and fit desired formatting. We use different models, prompting, sampling, filtering and editing for different data pipelines that compose the final data mixes.
We picked our final training recipes (data mix and hyperparameters) based on a battery of mostly internal automatic metrics. Both Jamba-1.5 models are fine-tuned with the same control tokens and formatting template, which we provide as a part of our release as a HF-compatible tokenizer and chat template; see the model card for details.
We give several notable examples of synthetic data generation:
We generate tabular data and accompanying question-answer pairs, as demonstrated in our work on table understanding . We then convert the tables into natural language paragraphs using a language model. Our generated training examples include extraction, aggregation, and attribution tasks vis-a-vis text corresponding to specific rows or columns in a given table.
Given a document, we prompt a language model to generate question-answer pairs, for both single and multiple paragraphs. We sometimes embed these examples within longer context by adding similar texts, to encourage long-context understanding with attribution.
We use the open-source Glaive function-calling dataset as a starting point, filtered with various heuristics and validations on the output schemas. To support parallel function calling, we first generate multiple valid parameter assignments for each function in Glaive. Next, we sample subsets of these valid parameter assignments, for the same function and across different functions, to generate user requests corresponding to the set of function calls. Finally, we prompt a function-calling language model to respond to these generated user requests and retaineonly responses where the function calls matched the original parameter assignments.
We defined a set of instructions that can be easily validated and synthesized prompts that include a generic document-drafting task with one or more constraints added to it. We generated completions for these prompts from a language model and used rejection sampling based on the validations of our fine-grained instructions plus a general-purpose reward model. To support instructions in system messages, we chose multiple prompts of this kind that share a fine-grained instruction instance and reformatted these prompts into a multi-turn conversation, with the instruction moved to the system message.
4 Some Observations
We share a few observations from the development of Jamba-1.5. While these are not fully explored, we hope they would inspire the community to look further into these issues.
First, while we included only a very small fraction of non-english data, for a few languages and only for specific skills in the post-training phase, our Jamba-1.5 models perform quite well in multiple languages. We did include multilingual data in the pre-training phase, as mentioned above. Thus we speculate that the models are able to use the learned knowledge from that phase when being post-trained mostly in English.
Second, our efficient Jamba architecture lowers the cost of fine-tuning on long contexts, allowing us to experiment more with a given budget. Thus we could experiment with multiple different training recipes at the post-training stage.
Finally, while preference tuning algorithms like PPO or DPO improve alignment between model outputs and human intent, we found that the combination of careful synthetic data generation, data filtering, and supervised fine-tuning is crucial for obtaining a strong post-trained model.
Evaluation
While we believe benchmarks are only partly correlated with success of real applications and user satisfaction, we report results on key public benchmarks. First, we report results on standard academic benchmarks. Then, we evaluate the model on chatbot benchmarks. Finally, we evaluate Jamba-1.5-Large on several long-context evaluations and a multilingual evaluation.
We compare with recent open-weight models of the same size range: LLaMA-3.1 70B and Mistral-Large-2-123B when comparing with Jamba-1.5-Large; LLaMA-3.1-8B and Gemma-2-9B when comparing with Jamba-1.5-Mini.
We report results with a wide range of standard academic benchmarks: MMLU , MMLU-Pro , GPQA , ARC-Challence , BBH , and HumanEval . We also evaluate on the IFEval instruction following dataset and the BFCL v1 function calling dataset . Finally, we report safety evaluations on RealToxicity and TruthfulQA .
Table 2 compares Jamba-1.5-Large to several publicly available models at similar sizes. All results are either taken from official sources or evaluated by us, as indicated in the table.In two cases we failed to obtain good results: Mistral-Large-2 fails to obtain good scores on ARC-C despite multiple attempts. LLaMA-3.1 models perform poorly on GSM8K with the standard strict evaluation mode, so we also report for them a flexible evaluation, which allows higher results. We observe that the Jamba-1.5 models perform similarly to recent state-of-the-art publicly available models on standard academic benchmarks, including knowledge, reasoning, instruction following and function calling capabilities. We also observe similar safety metrics as those reported in the literature. We refer to Section 7 for more information about our general approach for safety and alignment of models.
Importantly, the Jamba-1.5 models achieve these results while providing much better throughput and latency, as discussed above.
2 ChatBot Evaluations
In this section we evaluate the Jamba-1.5 models on two chatbot scenarios: Arena-Hard , a set of 500 challenging user queries that uses GPT4-Turbo as a judge, and WildBench , which uses GPT4-Turbo as a judge with a length bias mitigation. As Table 3 shows, Jamba-1.5 models obtain excellent reuslts in these evaluations, with Jamba-1.5-Large surpassing LLaMA-3.1 70B, but somewhat trailing behind Mistral-Large-2 123B, which has about 30% more active parameters.
3 Long-Context Evaluations
The released model handles context lengths of up to 256K tokens. In this section, we evaluate it on synthetic and naturalistic benchmarks that test its long-context capabilities.
We evaluate on the RULER benchmark, a set of 13 synthetic tasks aimed to assess long-context capabilities of language models. RULER includes 8 variants of needle-in-a-haystack retrieval tasks , including multiple ‘needles’ . It also has one variable tracking task where a chain of variable bindings should be returned, two aggregation tasks where one needs to return the most common words, and two question-answering tasks, where paragraphs cotraining answers from naturalistic datasets are inserted into random paragraphs to simulate long contexts.
The results are shown in Table 4. Among all publicly available and proprietary models, Jamba-1.5-Mini and Jamba-1.5-Large are the only ones with a confirmed effective length of 256K tokens. Gemini-pro reports good results up to 128K on the original RULER paper. However, we were unable to reproduce these results despite much effort. We examined Gemini-pro generations and noticed the model often fails to answer or generates a refusal. Since the official RULER results are from a preview version, we hypothesize that Gemini-pro had since undergone through updates that have hurt its performacne on RULER.
3.2 Infinite-Bench
Next we evaluate on Bench, a dataset designed to evaluate long-context abilities of language models, with an average length of 100K tokens. We focus on two English tasks on understanding long novels: question answering (EN.QA) and multiple-choice question answering (EN.MC). As Table 5 shows, Jamba-1.5 models perform very well in this case, outperforming similarly sized LLaMA-3.1 and Mistral-Large-2 models. (We do not report results with Gemma-2 9B due to its short context window of 8K.)
4 Multilingual capabilities
We perform a basic evaluation of Jamba-1.5 abilities in non-English langauges. In particular, we report results on the multilingual MMLU dataset as distributed through the LM Evaluation Harness . Table 6 shows the results, where Jamba-1.5-Mini performs similarly or better than its points of comparison. Jamba-1.5-Large is slightly behind its comparable models, but still exhibits good multilingual capabilities.
Alignment and Safety Considerations
Our approach to alignment of our models is driven by creating transparency between model behavior and customer expectations. Our models default to a business code of conduct based on our participation in industry standards bodies, think tanks and direct experience with our customers. We see this as an ongoing and evolving collaboration. In addition, companies have multiple ways to control model conduct to reflect their individual values and cultures such as additional training and fine tuning, system messages and prompt engineering. Overall, our AI code of conduct is based on the following objectives:
Align model behavior and output with company values and normative business decorum.
Clearly state tenets of intended behavior such that errors/bugs are easily discerned.
Collaborate with Customers and map behavior to their best practices.
Continuously gather feedback to monitor and actively improve behavior.
In line with our role in an OECD task force to develop a monitoring mechanism for applying the G7 Hiroshima Code of Conduct for Organisations Developing Advanced AI Systems, we have organized our model alignment work with the OECD values-based AI principles:https://oecd.ai/en/ai-principles inclusive growth, sustainable development and well-being; human-centered values and fairness; transparency and explainability; robustness, security and safety; and accountability.
For each of the first four principles we have detailed behavioral expectations or tenets and examples that can be used to train/align and test for compliance. The principle of accountability is focused on AI21’s role in taking responsibility for the behavior of the models. We submit that this accountability is demonstrated primarily through transparency and engagement with customers, regulators and independent 3rd-parties. Our engagement with OECD, Stanford’s HELM and FMTI and documents like this demonstrate this commitment, as well as our high ranking on the FMTI (2nd as of May 2024).
In total, we have created 60 tenets that map to the OECD principles. These tenets are stated as directives of behavior for our models to avoid. The full list will be made publicly available.
Conclusion
We have presented Jamba-1.5-Large and Jamba-1.5-Mini, two new large-scale models based on the Jamba hybrid Transformer-Mamba architecture. Both models achieve excellent performance in academic benchmarks, chatbot evaluations, and long-context evaluations, while offering improved latency and throughput, especially for long contexts. We release the model weights for use by the community in hopes that others build on this technology.
Contributions
Pre- and Post-Training Alan Arazi Barak LenzProject leads Chen Almagor Dan Padnos* Daniel Gissin* Daniel Jannai Dor Muhlgay Edden M Gerber Erez Safahi Gal Cohen Gal Shachaf Hofit Bata Inbal Magar Itay Dalmedigos Jhonathan Osin* Matan Danos Michael Gokhman Nir Ratner Noam Gat Noam Rozen Omer Antverg Omri Abend Opher Lieber* Orit Cohavi Raz Alon Shaked Meirom Tom Braude Uriya Pumerantz Yonatan Belinkov Yuval Globerson Yuval Peleg Levy
Serving & Infrastructure Amir Bergman Avshalom Manevich Barak Peleg Elad Dolev Eran Krakovsky Erez Schwartz Haim Rozenblum Mor Zusman Oded Fried Roman Glozman Shahar Lev Tomer Asida Yehoshua Cohen
Data Ben Aviram Dor Zimberg Ido Blass Ohad Leshno Rom Gilad Tom Ben Gal
Evaluation Clara Fridman Julie Fadlon Maria Rozman Naama Gidron Ro’i Belson Tal Ness
Project & Product Management Or Dagan* Roi Cohen Shaked Meirom* Tal Delbari Yoav Shoham