Focused Transformer: Contrastive Training for Context Scaling

Szymon Tworkowski, Konrad Staniszewski, Mikołaj Pacek, Yuhuai Wu, Henryk Michalewski, Piotr Miłoś

Introduction

Language models have served as a catalyst for substantial advancements in several areas, including natural language processing (Radford et al., 2019; Brown et al., 2020), code generation (Chen et al., 2021; Li et al., 2022), quantitative reasoning (Lewkowycz et al., 2022) and theorem proving (Polu and Sutskever, 2020; Jiang et al., 2022; Mikuła et al., 2023). One of the central challenges with language models is the effective incorporation of extensive new knowledge. The common practice of fine-tuning the model is not only resource-intensive and complex to manage, but it also does not always clearly indicate how to incorporate new knowledge. For example, fine-tuning on a text such as “Alice in Wonderland” does not equip the model to answer questions about the story itself, but rather it trains the model to predict the next token or complete masked sentences. A promising alternative – integrating the new knowledge within the context – doesn’t require training but is considerably restricted by the model’s effective context length. For this method to work with large knowledge databases (like large code repositories), the model needs to manage a context length extending to millions of tokens.

In this research, we highlight one of the primary obstacles in augmenting the context length: as the number of documents increases, the ratio of pertinent to irrelevant tokens diminishes. The standard training procedure frequently results in overlaps between keys connected with irrelevant values and those related to relevant ones, exacerbating the model’s task of differentiating between them. We term this challenge the distraction issue.

We propose the Focused Transformer (FoT), an innovative technique developed explicitly to address this issue. The Focused Transformer permits a subset of attention layers to access an additional context of (key, value) pairs through the k-nearest neighbors (kNN) algorithm, akin to the method used in (Wu et al., 2022). This mechanism effectively extends the total context length. The distinctive aspect of the Focused Transformer is its training procedure, drawing from contrastive learning. This method addresses the distraction issue and facilitates larger context capacities. Specifically, during the training phase, we deliberately expose the chosen subset of attention layers to both relevant and irrelevant keys (like negative samples from unrelated documents). This strategy incentives the model to differentiate keys connected with semantically diverse values, thereby enhancing their structure.

We introduce and make available LongLLaMAs (), fine-tuned OpenLLaMA models with FoT, demonstrating that our method does not require long context during training and can be applied to existing models. Notably, LongLLaMAs show significant improvements on tasks necessitating long-context modeling. In particular, they can manage a 256k256k context length on the passkey retrieval task (Mohtashami and Jaggi, 2023).

Our research contributions are the following:

1. We pinpoint the distraction issue as a significant challenge and a primary obstacle to scaling up the context length in Transformer models, particularly in multi-document scenarios.

2. We develop the Focused Transformer (FoT), designed to alleviate the distraction issue. FoT includes a unique training objective that improves the (key, value) structure, enabling the use of extensive additional context and k-nearest neighbors lookup to scale the context length.

3. Our method is simple to implement, and it provides the benefit of extending model context without modifying the architecture, facilitated by cost-effective fine-tuning. We demonstrate this on the 3B3B and 7B7B OpenLLaMA checkpoints. The resulting models, named LongLLaMAs, display enhancements on tasks that benefit from increasing the number of few-shot demonstrations in the extended context, such as TREC (Li and Roth, 2002; Hovy et al., 2001) and WebQS (Berant et al., 2013). We also prove that for passkey retrieval Mohtashami and Jaggi (2023), our LongLLaMA models successfully handle a 256k256k context length.

4. We further scrutinize FoT’s capabilities across various datasets and model sizes. We show that a FoT trained with a total context of 512512 tokens can extrapolate to 1616 million tokens in a benchmark dictionary lookup task. We also assess FoT on long-context language modeling tasks such as books (PG-19), mathematics (arXiv), code (GitHub), and formal proofs (Isabelle), where it exhibits improvements in perplexity over baselines.

Related work

A multitude of approaches have been developed to increase the context length of transformers, mostly focusing on alleviating the quadratic complexity of the attention computation. For instance, Transformer-XL (Dai et al., 2019) caches the previous context and enables the linear extension of context with the number of layers. Longformer (Beltagy et al., 2020) employs an attention mechanism that allows tokens to attend to distant tokens sparsely, reducing the computational complexity. BigBird (Zaheer et al., 2020), LongT5 (Guo et al., 2021), and (Dao et al., 2022) also use sparse attention to handle long sequences. Different efficiency considerations have been studied in (Kaddour et al., 2023), showing that they lead to limited gains. Hierarchical transformers (Nawrot et al., 2021, 2023) downsample activations in intermediate layers to reduce computation and enable longer contexts. COLT5 (Ainslie et al., 2023) proposes conditional computation to save memory and enable larger contexts. Memorizing Transformer (Wu et al., 2022) uses kNN lookup to pick up the most relevant tokens, which might also be seen as a way to reduce the computational complexity of attention. Our work adheres to this approach and aims to train a key space that handles longer attention context length (e.g., by mitigating the distraction issue) and, thus, has better long-context capabilities.

Fine-tuning LLMs for longer retrieval

Prior works such as RETRO (Borgeaud et al., 2022) (RETROfitting) and Memorizing Transformer (Wu et al., 2022) have demonstrated a promising path for fine-tuning existing LMs to add new capabilities without the need to retrain the entire model. In contrast to those approaches our method is not framed as a retrieval but as a way of extending the context of the model. In contrast to RETRO, we propose a single-stage method for context extension instead of a two-stage retrieve-then-embed approach. We provide a more detailed comparison with the Memorizing Transformer in Appendix C.3. More recently, a number of works have explored fine-tuning LLaMA to extend its context length. Landmark attention (Mohtashami and Jaggi, 2023) proposes a compression scheme of LLM’s context into landmarks, increasing the context length of LLaMA-7B to 32K32K. Position Interpolation (PI, (Chen et al., 2023) and (kaiokendev, 2023)) introduces a modification to the rotary positional encoding scheme that enables fine-tuning for 32K32K context. In contrast to this work, our method does not rely on positional encodings, following the findings from (Haviv et al., 2022). Removing positional encoding in additional context allows us to extrapolate to 256k256k tokens, although the model was only trained on sequences up to 8K8K, yielding theoretically unbounded context length.

Zero-shot methods

KNN-LM (Khandelwal et al., 2019) shows that one can improve the performance of a LLM by combining two probability distributions. One created by a pre-trained model, and one based on the similarity between the embedding of the currently processed token and the embeddings of tokens retrieved from a large database. Meanwhile, we extend the model context in a subset of attention layers, potentially allowing for reasoning within this extended context. Parallel Context Windows for Large Language Models (Ratner et al., 2023) introduces a method for extending the context of language models without training. They achieve this by embedding several context windows independently in parallel and allowing only a subset of tokens to attend to all windows. On the other hand, we fine-tune existing models and allow all tokens to attend to all previous tokens but only in a subset of layers. Additionally, our method allows us to improve the structure of the key-value space of the existing models.

Contrastive learning

Contrastive learning aims to learn good representations by comparing positive and negative examples. CLIP (Radford et al., 2021) and SimCLR (Chen et al., 2020) are two popular contrastive learning methods that have achieved state-of-the-art performance in the image domain. During contrastive pre-training, negative examples are kept in the same batch to learn to distinguish them from positive examples. Scaling the batch size in contrastive learning has been demonstrated to enhance the quality of representations, as shown in (Gao et al., 2021b). It has been suggested (Gao et al., 2019) that the embedding space in language modeling suffers from degeneracy, where embeddings are tightly packed in a narrow cone, making it difficult to distinguish between them. TRIME (Zhong et al., 2022) proposes a training approach designed for training LMs with memory augmentation, which uses negatives to improve the quality of representations. The main difference between this and our approach is that we incorporate negatives into the chosen subset of attention layers instead of interpolating in the output layer and use the standard language modeling loss. TRIME (Zhong et al., 2022) also focuses on retrieval from large databases, whereas we focus on extending the context of the model. ContraCLM (Jain et al., 2023) applies contrastive losses at both the token and sequence levels during training to promote more uniformly distributed, isotropic representations. It is shown to enhance the discrimination of representations on textual semantic similarity benchmarks. While ContraCLM focuses on improving the general expressiveness of representations, our work introduces contrastive-inspired techniques designed specifically for training the attention mechanism to handle longer context lengths. Nonetheless, exploring other contrastive learning objectives could be beneficial for further improving the key structure in future work.

FoT: Focused Transformer

Our method, the Focused Transformer (FoT), is a simple plug-and-play extension of transformer models and can be used both to train new models or fine-tune existing, possibly large, models with longer context. To this end, FoT uses memory attention layers and the crossbatch training procedure. Memory attention layers enable the model to retrieve information from the additional context at inference time, effectively extending the context. The crossbatch training procedure biases the model to learn (key,value)(key,value) representations, which are easy to use by a memory attention layer. See Figure 2 for an overview of the FoT architecture and Appendix L for pseudocode.

2 Crossbatch training procedure

The operation is fully differentiable, and thus, we improve all the (key,value)(key,value) pairs in pδp^{\delta}. Two, the procedure is easy to implement; it does not require any additional loss (i.e., uses the standard transformer training objective) and is done on the level of the data loading pipeline and a minor self-attention change. The only new hyperparameter is dd, which prescribes the ratio of positive to negative samples. Typically, we find it beneficial to start with small d≤8d\leq 8 (otherwise, the model tends to ignore the previous local context) and later switch to bigger values, say d≥64d\geq 64. Appendix B.3 provides more details about the method. Listing 1 outlines an implementation of the crossbatch.

3 The distraction issue

In this section, we conceptualize what we call the distraction issue and hypothesize it is one of the key problems in dealing with long multi-document contexts (like large code repositories). Namely, during the standard training, the model is not incentivized to distinguish the keys from different documents. We measure that the attention mass is evenly spread on the related and unrelated documents; see Figure 3. More precisely, for a document δ\delta, let wijw_{ij} be the softmax weights related to pijδp^{\delta}_{ij} constructed as described in Section 3.2. We define the positive attention mass as rd:=∑jw1j/∑i=1d∑jwijr_{d}:=\sum_{j}w_{1j}/\sum_{i=1}^{d}\sum_{j}w_{ij}. We observe that rd≈1/dr_{d}\approx 1/d, which can be interpreted as the fact that the attention is equally distracted by the positive (coming from the current document at i=1i=1) and negative keys. This is an undesirable property since when scaling the memory, the attention becomes increasingly distracted. We show that the crossbatch mostly alleviates the distraction issue, resulting in a focused attention. More information can be found in Appendix B.4. In Section 5.3, we also show that the distraction issue has a harmful effect on metrics like perplexity.

LongLLaMA : extending LLaMA’s context length with FoT

One of the promises of our work is that FoT can be used to fine-tune already existing large models to extend their context length. In this section, we show that this is indeed the case. We use OpenLLaMA-3B and OpenLLaMA-7B models trained for 1T1T tokens as starting points and fine-tune them with FoT. We show that the resulting models, which we call LongLLaMAs, are capable of extrapolating beyond their training context length (even up to 256K256K) and retain the performance on short-context tasks. We release the inference code on GitHub: https://github.com/CStanKonrad/long_llama and the LongLLaMA-3B checkpoint on Hugging Face: https://huggingface.co/syzymon/long_llama_3b. We note that our checkpoint is backward compatible, i.e. can be used with any existing LLaMA inference code (both in Hugging Face and other implementations), albeit without long-context capabilities.

The architecture of the models is the same as OpenLLaMAs, see Geng and Liu (2023) and Appendix A.1. We use L={6,12,18}\mathcal{L}=\{6,12,18\} (resp. L={8,16,24}\mathcal{L}=\{8,16,24\}) as the memory layers for 3B3B (resp. 7B7B) LongLLaMA model. We fine-tune the models on 10B10B (resp. 3B3B) tokens using FoT, 8k8k context length and our dataset mixture based on RedPajama (TogetherComputer, 2023), see Appendix A.3.

There are three minor differences from the standard FoT procedure. First, we retain the positional encodings in the local context of the memory layers (this is not necessary for FoT, but makes our checkpoints fully compatible with any existing LLaMA inference codebase). To be more precise, queries and keys from the local context (up to 2K2K tokens) receive the standard LLaMA rotary positional encoding, whereas memory keys are encoded as if they had position 0 in the local context window. Second, we use dense attention instead of the kNN retrieval, as we found only marginal performance differences, and it is simpler to implement. Third, we modify the crossbatch training procedure to have more fine-grained control over the number of additional contexts and the ratio of positive to negative samples. All these differences are detailed in Appendix A.2.

2 Context length extrapolation on the passkey retrieval task

We first measure the effective context length of LongLLaMA, namely the distance for which tokens can effectively attend each other. We use passkey retrieval introduced in (Mohtashami and Jaggi, 2023), a synthetic task designed to measure this property. In this task, the model has to retrieve a passkey placed randomly in a long prompt. Results are shown in Figure 1 - importantly, our 3B3B model is capable of solving this task much beyond its training context length 8K8K, achieving 94.5%94.5\% accuracy for prompts of length 100k100k and 73%73\% for 256k256k.

3 Question answering over research papers

In Table 6 we present the performance on the validation set of Qasper (Dasigi et al., 2021) from SCROLLS (Shaham et al., 2022) and compare our results to LongChat 7B (Ma and Zhang, 2023) and two baseline short-context models. We note that our model shows gains from increased context length.

4 Improving few-shot learning accuracy with longer context

We measure long-context capabilities of these models on two downstream tasks, TREC question classification (Li and Roth, 2002; Hovy et al., 2001) and WebQS question answering (Berant et al., 2013). We follow the experimental setup of (Hao et al., 2022). Namely, we few-shot prompt the models with as many demonstration examples as possible up to the given context length. We do not use structured prompting like in (Hao et al., 2022) - instead, we directly provide all demonstrations in context.

We observe significant accuracy gains from longer contexts on TREC and some improvements on WebQS (see Table 1). The TREC dataset consists of 5050 classes. A model is tasked to predict the class label given in-context examples. Only 100100 examples fit the standard context length (2K2K); it is not unusual that no class example is present for a given question, making the task impossible. Increasing the context length and the number of examples mitigates this risk. Moreover, having more demonstrations of the given class is also likely to be beneficial.

5 Comparison to standard long-context fine-tuning

In this section, we compare FoT to standard long-context fine-tuning, showing that it already achieves better performance for the context length used for fine-tuning and, importantly, that it can extrapolate beyond this context length, which is not the case for the baseline.

For comparisons, we fine-tune two models, one trained with FoT and another one (baseline) with standard fine-tuning (done similarly to (MosaicML, 2023; Nijkamp et al., 2023)). In both cases, we use 3B3B models fine-tuned on 1B1B tokens using the 4K4K context length. We evaluate both models on a number of few-shot downstream tasks in the setting described in Section 4.4.

In most cases, see Table 2, we observe accuracy improvements when more few-shot demonstrations are provided in the extended context (from 2K2K used by OpenLLaMA to 4K4K used in our fine-tuning). On TREC, the gains from additional context are significant for both models, while on WebQS, the standard fine-tuning baseline does not provide any improvement from extended context. Notably, the model fine-tuned with FoT enjoys further accuracy gains when evaluated with context lengths beyond its training length (6K6K and 8K8K). This shows extrapolation capabilities of FoT, which are not present in the baseline (see e.g. Figure 1).

6 Performance on short-context tasks

Fine-tuning for longer contexts could hurt performance on the original context length (2K2K), as the training data distribution changes. We show that this is not the case for the LongLLaMA models by evaluating them using the LM Evaluation Harness library (Gao et al., 2021a). On most tasks, the performance is kept intact; see Appendix A.4 for details. This also confirms that LongLLaMAs could be used as a drop-in replacement of LLaMA models as they are compatible with the original LLaMA inference code.

Analysis of FoT

In this section, we perform extensive experiments on smaller models to analyze and further validate our approach. In particular, we answer the following questions: (1) How does FoT perform when scaling the context length at inference time? (2) Can FoT be used to extend the context length of an existing, pre-trained model? (3) How effectively can it handle distractions, and how does this capability translate to enhanced performance in long-context language modeling tasks? Moreover, we provide ablation studies of our method and additional analysis.

Evaluation We distinguish two evaluation settings: single-document (abbreviated to single-doc) and multi-document (abbreviated to multi-doc). The single-doc setting is typically used for evaluating models that process long contexts. Here, we clear the memory for each new document, ensuring that only the current document is available in the context. The multi-doc setting retains memory across multiple documents without resets. This scenario tests whether the model can ignore irrelevant information and focus on the relevant data, which can be useful in setups like repository-level code generation.

Datasets We evaluate on the following long-context language modeling datasets: PG-19 (English books), arXiv (mathematical papers), GitHub (code), and Isabelle (formal proofs). PG-19 (Rae et al., 2019) is a large dataset of English-language books published prior to 1919, sourced from the Project Gutenberg archive. This dataset is a well-established benchmark for evaluating long-context language models (Sun et al., 2021). The arXiv dataset contains LaTeX source of papers labeled as "Mathematics" that were obtained by downloading articles through the arXiv Bulk Data Access. The token count per paper in this dataset is comparable to that of a book in PG19. For details on the remaining datasets, refer to Appendix H.

2 FoT fine-tuning and context length extrapolation

FoT is a minimal modification to the standard transformer architecture; therefore, it is possible to fine-tune existing models to endow them with a longer context length via the memory attention layer, as we already demonstrated in Section 4. In this section, we deepen this analysis (on a smaller model) by studying perplexity improvements on various datasets.

As a base model, we use a standard transformer model pre-trained for 100k100k steps with context of 1K1K tokens using the standard objective and fine-tune with the FoT objective (i.e. crossbatch). The data used for both fine-tuning and pre-training is the C4 dataset Raffel et al. (2019a) (we omit documents shorter than 2K2K tokens). The fine-tuning phase takes 10k10k steps. We use the crossbatch dimension d=128d=128 and local context of 1K1K tokens (context is 2K2K during training). We evaluate models in a zero-shot way on 44 language modeling datasets, which require long context: arXiv, PG-19, GitHub and Isabelle, see Section 5.1 and Appendix E for details.

In Table 3, we observe that FoT enjoys steady perplexity gains up to 64K64K tokens, although it was fine-tuned only with the 2K2K total differentiable context length. We compare the model perplexity to the following baselines: Memorizing Transformer (MT) (Wu et al., 2022) fine-tuned with the local context of 1K1K and memory size of 16K16K, and Transformer-XL (Dai et al., 2019) fine-tuned with both local context and window length of 1K1K. To ensure a fair comparison, all three models are fine-tuned from the same base checkpoint. When evaluated with a context of 2K2K, our method achieves results on par with the Transformer-XL baseline, which has access to the previous context in all layers, unlike MT and FoT. Compared to the MT baseline, we achieve better scaling when evaluated with 64K64K context length and significantly better perplexity values. Unlike MT, our method does not require training on long sequences, which is reflected by the lower perplexities of FoT when evaluated in the zero-shot setting. For more details, see Appendix G.

We also confirm the context extrapolation abilities using a synthetic dictionary lookup task. In this task, the model is first provided with ki:vik_{i}:v_{i} mappings and then asked what value is associated with a particular key. We train 3737M parameter models using documents of length 512512. Figure 10 shows that FoT, after 55k steps of training, can effectively utilize memory consisting of 1616M tokens achieving accuracy above 92%92\%. Details can be found in Appendix F.

3 Handling distractions in language modeling tasks

In this section, we measure how handling distractions in the multi-document setting helps in language modeling. We pick the PG-19 dataset (Rae et al., 2019) and measure the perplexity of the next token prediction (language modeling task) when varying the size of multi-doc memory (in this case consisting of books). Intuitively, the memory tokens corresponding to the current book might be beneficial (which is also confirmed in (Wu et al., 2022)), while the ones from the other books are unlikely to be useful and thus are distractions.

We observe, see Figure 9, that higher values of the crossbatch dimension dd lead to better perplexity. This aligns with the observations in Section 3.3, indicating that by mitigating the distraction issue, we experience benefits in language modeling.

Moreover, all versions of FoT are able to utilize memory and achieve much better perplexity than the standard Transformer (no memory). Unsurprisingly, perplexity increases with memory size, but we stress that this happens gracefully. In the standard variant of FoT (bold line), the perplexity increases only by 0.180.18 when scaling to >500k>500k tokens. Importantly, the perplexity of FoT is close to this of Memorizing Transformer with the single-doc memory, which we treat as a soft lower bound since it is not exposed to distractions from unrelated books.

4 Context length extrapolation in single-doc

The original motivation behind FoT is to improve the multi-doc setting performance by handling distractions. Interestingly, our method also helps to extrapolate to longer contexts, even when evaluated in the single-doc setting.

To study this, we perform FoT fine-tuning (as in Section 5.2) and evaluate the perplexity of the resulting model on the PG-19 dataset with different context lengths in the zero-shot fashion. To deepen the analysis, we introduce an additional parameter ww (the number of previous contexts used in cross batch training procedure). We provide results for w=1w=1 (the standard setting for FoT, that corresponds to the total differentiable context being 2⋅10242\cdot 1024) and w=2w=2 (corresponding to the total differentiable context 3⋅10243\cdot 1024).

We observe, see Figure 9, improvements when context grows, even far beyond the training context length, which reaffirms the hypothesis that FoT helps with extrapolation to longer contexts. Moreover, d=2d=2 is significantly better than d=1d=1. When comparing d=1d=1 and w=2w=2 to d=2d=2 and w=1w=1, we observe that the former is slightly better. This is natural, as the former has longer training context.

5 Ablations and design choices

In Appendix C we present ablations on our design choices. In particular, we note the importance of differentiability and the inclusion of negatives. We also discuss the relation to Memorizing Transformer. We note that due to the limited resources we have followed the Memorizing Transformer in the choice of memory layers.

Limitations and future work

Our research opens a few avenues for future work. We list them as well as challenges and limitations.

Scaling up context This is by far the most important future research direction. The challenges start from purely engineering, storing more than 1616M (key,value)(key,value) pairs will require a distributed multi-node system. In our experiments, we use the exact kNN search, which is not scalable to large memory. Using approximate kNN search will require a lot of engineering effort, as well as careful evaluation of the impact of the approximation on the model performance.

Scaling up crossbatch We observed that increasing dd is beneficial. In our experiments, we used d=64d=64 or d=128d=128, which is the maximum value that fits into the memory of a single TPUv3/TPUv2 machine, see also Appendix I. In future work, we want to further increase dd as well as test on devices with bigger memory or utilize multi-node training. We also note that crossbatch increases the training cost, but only in a subset of layers.

Exploring contrastive learning The FoT training is inspired by rather basic contrastive learning (CL) techniques. We show that this improves the key structure so that the distraction issue is mitigated. We expect that other CL methods could be beneficial, for example, hard negative mining to utilize a larger memory during training (see (Lindgren et al., 2021)). We leave this for future work.

Combining with other methods Developing long-context methods is an active research field, see Section 2. We believe that some of these methods could be combined with FoT, resulting in mutually beneficial interactions.

Acknowledgments and Disclosure of Funding

We gratefully acknowledge the TPU Research Cloud program, which was instrumental to our research by providing significant computational resources. Parts of the project were realized using the resources of Poznańskie Centrum Superkomputerowo - Sieciowe. We would also like to thank Markus Rabe for reviewing the initial manuscript and Christian Szegedy, Charles Staats, and DeLesley Hutchins for helpful discussions. We are also grateful to Xinyang Geng and Hao Liu for releasing OpenLLaMA checkpoints and the EasyLM library (Geng, 2023), allowing for training these models, which significantly accelerated our research. Piotr Milos was supported by the Polish National Science Centre grant 2019/35/O/ST6/03464. Henryk Michalewski was supported by the Polish National Science Center grant UMO-2018/29/B/ST6/02959.

References

Broader Impact

Recent rapid developments in language models have brought a lot of new capabilities. At the same, these raised concerns about the social impact and very animated discussions in the community. Our work develops a generic technique, which in principle, could be applied to virtually any language model and thus, by extending their capabilities, exacerbate threats. We note, however, that FoT does not create any new threats. Thus, we refer to the existing body of knowledge on the broader impact of language models, see e.g. Borgeaud et al. .

Appendix A LongLLaMA

OpenLLaMA [Geng and Liu, 2023] is an open-source reproduction of LLaMA [Touvron et al., 2023]. It uses a decoder-only architecture with rotary positional embeddings, and a few changes including pre-normalization with RMSNorm [Zhang and Sennrich, 2019], and SiLU activation [Elfwing et al., 2017]. A SentencePiece [Kudo and Richardson, 2018] tokenizer with 32k vocabulary size is used.

A.2 Extending context length with FoT

Positional encodings To achieve backward compatibility with the original LLaMA, we retain positional encodings in the local context. The tokens outside the local context are assigned the same position as the first token in the local context.

Dense attention to longer context To make the implementation simpler and less dependent on external software, we resign from using kNN lookup and perform attention over the whole memory. We have found only marginal performance differences between those two approaches to memory attention.

Crossbatch details For the 3B LongLLaMA model, we set L={6,12,18}\mathcal{L}=\{6,12,18\} as the memory layers. We vary the number of additional contexts d∈{0,2,3}d\in\{0,2,3\} across elements of the batch by dividing batch entries into four segments of equal size. Elements from the first segment only see local context (d=0d=0). Elements from the second segment see two additional contexts (d=2d=2), one from the same document (positive) and one from a different one (negative). Elements from the third segment see three additional contexts, two positives, and one negative. The last segment consists of elements exposed to three additional contexts coming from the same document. We abbreviate this setup as 14(0,0),14(1,1),14(2,1),14(3,0)\frac{1}{4}(0,0),\frac{1}{4}(1,1),\frac{1}{4}(2,1),\frac{1}{4}(3,0).

For the 7B LongLLaMA model, we set L={8,16,24}\mathcal{L}=\{8,16,24\} as the memory layers. Here we divide batch entries into four segments and use the following setup: 14(0,0),14(1,2),14(2,5),14(3,4)\frac{1}{4}(0,0),\frac{1}{4}(1,2),\frac{1}{4}(2,5),\frac{1}{4}(3,4).

A.3 LLaMA fine-tuning dataset

We use a mixture based on RedPajama [TogetherComputer, 2023] and The Stack [Kocetkov et al., 2022] with the following proportions of each subset:

All subsets apart from python are taken directly from RedPajama. For the python subset, we gather Python source code from The Stack and, to obtain long documents for training, concatenate files that are in the same subdirectory in random order, using a similar procedure as for the GitHub dataset in Section H. Additionally, we filter out short documents for some subsets of the original RedPajama, namely shorter than the Min. doc. length column indicates.

In case one document is too short to span across several contexts for crossbatch, then we concatenate it with the next document from the dataset.

A.4 Language Model Evaluation Harness

To ensure that the performance of LongLLaMAs has not degraded in short context scenarios, we evaluate our models on the Language Model Evaluation Harness benchmark [Gao et al., 2021a]. Table 5 compares our results with OpenLLaMA [Geng and Liu, 2023]. Similarly to the authors of OpenLLaMA, we omit CB and WSC tasks.

A.5 Question answering over research papers

We evaluate the context utilization of our model on the validation set of Qasper [Dasigi et al., 2021] from SCROLLS [Shaham et al., 2022]. Details are in the Table 6.

Appendix B Architecture

This section describes the architecture and crossbatch details for non-LLaMA-based models presented in this paper. The main differences are that for LLaMA-based models (LongLLaMA) we maintain the positional encodings (with a slight modification detailed in A.2), do not introduce the attention temperature parameter, and replace kNN with full dense attention.

For non-LLaMA-based models we use the transformer architecture introduced in [Vaswani et al., 2017] with a few standard changes. First, we use only the decoder without the encoder part. Secondly, we perform layer normalization before the input of both the attention and feed-forward modules. Additionally, we use Rotary Position Embedding [Su et al., 2021], normalize keys and queries [Henry et al., 2020], and introduce a learnable temperature parameter for each attention head.

The hyperparameters for each model size can be found in Appendix E. For training the models on PG-19, we use the standard T5 tokenizer with 3232k vocabulary [Raffel et al., 2019b]. The larger models in Section 5.2 are trained with a custom SentencePiece tokenizer [Kudo and Richardson, 2018] with 6464k vocabulary size.

B.2 Memory attention layer

where s(key)s(key) is the softmax score for keykey. This softmax is calculated as follows:

where τ\tau is a temperature parameter. In this approach, we do not distinguish between the local context and the memory.

Another way of integrating MtopM_{top} is via gating. In this approach, we separately compute the attention value vMv_{M} for MtopM_{top} and for the local context vCv_{C} (using the standard Transformer formula). Then we use a gating mechanism to combine them:

where σ\sigma is the sigmoid function and bgb_{g} is a trainable bias. The gating approach was proposed in [Wu et al., 2022], see formula [Wu et al., 2022, (2)].

We found our approach, i.e. using (1), to be equally effective, see Figure 4. At the same time, (1) is simpler and does not require additional parameters. Thus, we use it in our experiments.

We do not use the τ\tau parameter for LongLLaMAs as their architecture does not normalize keys and queries. For LongLLaMAs, we also replace the kNN search with dense attention and retain positional encodings (see Appendix A.2).

B.3 Crossbatch training procedure

Note that the only difference between (1) and (2) is the source of the additional (key,value)(key,value) pairs: pδp^{\delta}. This, in particular, implies that all the operations with respect to the previous context are differentiable.

The number of different documents is equal to bSb_{S} (the batch size, i.e. each document has a separate index in the batch). Assume that document δ\delta has index ii. We include into pδp^{\delta} all tokens from CprevC_{prev} with the batch indices in {i,(i+1)mod  bs,…,(i+d−1)mod  bs}\{i,(i+1)\mod b_{s},\ldots,(i+d-1)\mod b_{s}\}.

B.4 Qualitative analysis

Table 7 provides a brief qualitative analysis of FoT. It shows that the model can handle distractions and retrieve the parts of the character name from the multi-document memory in the PG-19 task dataset and appropriate definitions from the large dictionary (dictionary lookup task).

B.5 Memorizing Transformer

The Focused Transformer shares many similarities with the Memorizing Transformer [Wu et al., 2022]. In this section, we summarize the differences between those two models.

The key difference between these two methods lies in the training procedure. Our method uses crossbatch, see Section B.3, which, in a nutshell, is the standard transformer training objective, but we additionally attend to tokens from the previous context window, both from the same and different documents, see Appendix B.3 for details. The Memorizing Transformer trains on tokens retrieved from the same document (it was envisioned for single-doc).

FoT does not use memory during training, while MT does. This may result in faster training; moreover, FoT always uses the most up-to-date values, while MT uses the values from memory, which may be outdated.

FoT is differentiable through all (key,value)(key,value) pairs, while MT does not differentiate through the retrieved tokens. We argue that this is key for joint training of well-structured key, value, and query embeddings and, consequently, good model performance.

FoT does not require long documents in the training set, while MT does in order to capture long dependencies in memory. This is practically important, as many popular datasets consist of short documents.

We speculate that there may be benefits in blending these two approaches. One can, for example, argue that MT is better at providing ’hard’ negatives for the model. We provide a proof-of-concept experiment in Appendix C.3, and leave this for future work.

Inference

Both models use a very similar memory attention layer. The difference is how the retrieved (key,value)(key,value) pairs are integrated. FoT treats the retrieved information in the same way as the local context. MT uses a gating mechanism. Details are provided in Section B.2.

Appendix C Ablations

In this section, we focus on two key properties of crossbatch training procedure: differentiability and the inclusion of negatives. We also discuss the relation to Memorizing Transformer in terms of the training protocol and memory integration. We refer to Appendix B.5 for a detailed technical description of differences between FoT and Memorizing Transformer.

We compare FoT to Memorizing Transformer, which uses a non-differentiable memory of keys and values during training. In the multi-doc experiment presented in Figure 5, both MT and FoT are trained with local context of 512512. We observe that FoT is significantly better when the context is expanded during inference, which confirms that differentiable keys and values are beneficial.

We also check whether differentiable keys and values can improve the performance in the single-doc setting. For this, we compare FoT with d=1d=1 to MT with memory consisting of the previous local context. Table 8 confirms that differentiable keys and values can also help in this scenario.

C.2 Importance of negatives

We reaffirm the importance of negatives in a multi-document setting. In previous experiments in Figure 3, we already observed that increasing the number of negatives (i.e., increasing dd) results in more attention mass being dedicated to relevant tokens. In Figure 6, we additionally show that the lack of negatives in training (d=1d=1) results in a significant deterioration in model perplexity when the context length grows. This confirms that both using negatives and differentiability are important for FoT to work well.

C.3 Relation to Memorizing Transformer

Memorizing Transformer Wu et al. is closely related to our method. The two key differences are 1) the training protocol and 2) how the memory is integrated into the model. In this section, we provide additional insights into these differences.

Training protocol In the previous sections, we have discussed the benefits of the crossbatch training, namely using the contrastive-inspired objective and backpropagating through the previous context. A potential advantage of the MT approach is that it is exposed to the whole memory during training (instead of just the previous context). We performed a proof-of-concept experiment combining the two approaches to explore this further.

Namely, we trained the model for 499499k steps using crossbatch and fine-tuned it with the MT objective for 11k steps. Interestingly, we observed a significant improvement compared to the MT training with the same step budget, see Figure 7. We believe there is further room to explore various training protocols combining the best of both worlds.

Memory integration FoT uses a simple memory integration approach where the (key,value)(key,value) pairs retrieved by kNN lookup are treated the same way as the local context. In contrast, MT uses a gating mechanism, a weighted average of the memory, and local values; see details in Appendix B.2. We evaluated both approaches and found no difference in performance between these two memory integration methods. However, we decided to use our approach because it does not require any architectural changes (and thus makes fine-tuning existing models easy). For these reasons, we recommend using it. We speculate that the reason why the gating is not needed in FoT is another benefit of the fact that the crossbatch training backpropagates through the (key,value)(key,value) pairs from the previous context CprevC_{prev} in contrast to MT that cannot backpropagate there and needs to rely on local context when computing gradients for keys and values. Another reason might be the fact that CprevC_{prev} is embedded for each batch, and thus staleness (see [Wu et al., 2022, Section 3.2]) is avoided.

Appendix D Additional figures

Appendix E Hyperparameters

Table 9 shows hyperparameters used in our experiments. We used context length 512512 unless stated otherwise. In Appendix F, Section 5.3, Section 5.5, we use the total batch size of 32K32K tokens. In Section 5.2 and Section 5.4, the total batch size is 128K128K tokens.

For the experiments described in Section 5.3 and Section 5.5 we performed the following hyperparameter sweeps:

Batch size: {8K,16K,32K}\{8K,16K,32K\}, chosen: 32K32K.

For the dictionary lookup task (Appendix F) we checked the following hyperparameters:

Number of dictionary tokens in training step: {16K,32K}\{16K,32K\}, chosen: 32K32K. Note that during the training number of document tokens dedicated to the dictionary is the same as the number of tokens dedicated to questions.

For most of the other hyperparameter choices, we followed [Wu et al., 2022], to provide a fair comparison.

In Sections 5.3, 5.4 and 5.5 for models with d∈{1,2,4,8}d\in\{1,2,4,8\} we used constant schedule, and for models with d=64d=64 we trained with d=2d=2 for 450k450k steps and switched to d=64d=64 for the final 50k50k steps. In Appendix F we trained with d=1d=1 until the model reached 98%98\% accuracy and then we switched to d=128d=128. For the 184M184M model in Section 5.2, we randomly sampled dd from {2,128}\{2,128\} in each training step.

Appendix F Dictionary lookup task

We propose a dictionary lookup task to test whether the model trained using our crossbatch method can utilize a large memory database. Documents in this task are split into two parts. The first part defines keys and their associated values using the records of the format:

where is a special token that denotes the beginning of the defining sequence,

The second part consists of queries about the values associated with the previously defined keys. The queries are in the following format:

where is a special token that denotes the beginning of the query. We mask the loss so that for such a question, only v1,v2,v3,v4v_{1},v_{2},v_{3},v_{4} are included. We use a vocabulary of 6464 tokens, keys and values are described using 44 tokens.

During training, we use documents comprising 512512 tokens. The first half of each document consists of definitions, whereas the second one consists of questions. For FoT, we use a local context of 256256, thus the model needs to use the memory attention layer to answer the questions correctly. We start with d=1d=1 and increase to d=128d=128 as soon as the model is able to reach 98%98\% training accuracy. During the inference, we use k=32k=32 (the number of keys retrieved by kNN). As a baseline, we use a standard transformer model trained with the context length of 512512. In evaluation, we test different local context lengths, which quickly leads to very poor results.

In evaluation, we use longer documents but make only the last 256256 tokens correspond to questions. That is, as the context gets bigger (the token axis on Figure 10), the number of definitions increases, but the number of queries remains the same.

Appendix G FoT fine-tuning

For comparison in Table 3, our model is pre-trained for 100k100k steps with a total batch size of 128 (128K128K tokens per step, with 10241024 local context). Then we fine-tune both FoT and baselines for additional 10k10k steps with the same batch size. When fine-tuning FoT, we randomly sample dd from {2,128}\{2,128\} in each training step to prevent the model from overfitting to a large additional context length during training.

Appendix H Datasets

Section 5.1 outlines essential details concerning the PG-19 and arXiv datasets employed in this study. Now, we will present details about the remaining datasets:

We obtained a large corpus of permissively licensed Github repositories using BigQuery. By filtering for specific file extensions (C, C++, Java, Python, Go, and TypeScript), we captured individual source code files that are often short but have numerous dependencies and cross-references within the repository. To preserve the structure while shuffling the files and subdirectories in a random order, we concatenated all the files within each repository, treating subdirectories as a unit, similarly to Wu et al. .

Isabelle

The Isabelle corpus comprises formal mathematical proofs in the form of theories written in a formal language. We combined theories from The Archive of Formal Proofs (from October 2021) https://www.isa-afp.org and the Isabelle standard library to create a corpus of theories licensed as open source. Each theory focuses on topics like foundational logic, advanced analysis, algebra, or cryptography and consists of multiple files containing proofs. Similar to the GitHub corpus, the files within each theory are concatenated into a single document. However, unlike the Github corpus, we arrange the files based on their import dependencies, ensuring that later files can utilize sub-theorems proven in earlier files.

Appendix I Hardware and technical details

We used TPU virtual machines from the Google Cloud Platform (GCP). Each TPU virtual machine has 8 TPUv2 / TPUv3 cores totaling 6464GB / 128128GB of device memory, 96 CPU cores, and over 300GB of RAM. In larger-scale experiments (Section 5.2) we used machines with 32 TPUv3 cores. For training the LongLLaMA checkpoints, a TPUv3-128 pod provided by the TPU Research Cloud was used, which we gratefully acknowledge.

Appendix J Randomness

To evaluate the significance of our results, we conducted multiple runs for selected experiments in our study. In Figure 10, we calculate error bars, showing the minimum and maximum value over 10 runs of the same experiment. For the arXiv baseline experiment in Appendix K, we performed three runs with different random seeds and calculated their standard deviation, which is equal to 0.0020.002 perplexity. However, due to resource constraints, we were unable to conduct multiple runs for all experiments. Our preliminary findings indicate that the observed variance was minimal compared to the impact observed from other factors under investigation.

For the calculation of test perplexities, we used 1M1M tokens.

Appendix K Additional experimental results

This section presents additional empirical results, providing a detailed comparison of FoT with the Memorizing Transformer [Wu et al., 2022] baseline. Both models are trained for the same number of 500k500k steps with local context of 2K2K and evaluated on the arXiv dataset in the single-document setup, following [Wu et al., 2022]. In particular, we study how models trained with a given context length perform when evaluated with different context lengths. These experiments differ from those in Section 5.2, as the models were both trained and evaluated on the same dataset (arXiv), unlike the C4 training and zero-shot evaluation done in Section 5.2.

The MT baseline in Table 10 with a memory length of 2K2K struggles to utilize additional context beyond 32K32K tokens effectively. The model trained with 8K8K memory performs significantly better when evaluated with longer contexts, showing further perplexity gains at 64K64K tokens. We observe diminishing returns when scaling up the training memory length to 16K16K tokens and beyond.

Using the same setup, we study the performance of FoT while varying dd and ww configurations, similarly to Section 5.4, see Table 11. Parameter values w=1w=1 and w=2w=2 correspond to additional context lengths of 2K2K and 4K4K, respectively. In an apples-to-apples comparison to MT with 2K2K additional context length, FoT outperforms the MT baseline, which shows the importance of trainable keys and values (see also Section C.1). Moreover, we confirm the findings from Section 5.4 that d=2d=2 works significantly better than d=1d=1 in all settings. Our best configuration achieves 2.1482.148 perplexity with 4K4K additional context during training, compared to 2.1642.164 of MT with 16K16K additional context.

Appendix L Code

In Listing 2, we show the FoTs attention code (i.e., the code for the memory attention layers and crossbatch training), see Section 3, Appendix B.2, Appendix B.3. We note that the changes to the code are small; they are localized to the memory layer (the other layers follow the standard transformer protocol) and do not require any new trainable parameters.