How to train a Million Context LLM — with Mark Huang of Gradient.ai
About this episode
From the show’s notesAI Engineer World’s Fair in SF! Prices go up soon. Note that there are 4 tracks per day and dozens of workshops/expo sessions; the livestream will air the most stacked speaker list/AI expo floor of 2024. Apply for free/discounted Diversity Program and Scholarship tickets here. We hope to make this the definitive technical conference for ALL AI engineers. Exactly a year ago, we declared the Beginning of Context=Infinity when Mosaic made their breakthrough training an 84k token context MPT-7B.
A Brief History of Long Context Of course right when we released that episode, Anthropic fired the starting gun proper with the first 100k context window model from a frontier lab, spawning smol-developer and other explorations. In the last 6 months, the fight (and context lengths) has intensified another order of magnitude, kicking off the "Context Extension Campaigns" chapter of the Four Wars: * In October 2023, Claude's 100,000 token windows was still SOTA (we still use it for Latent Space’s show notes to this day). * On November 6th, OpenAI launched GPT-4 Turbo with 128k context. * On November 21st, Anthropic fired back extending Claude 2.1 to 200k tokens. * Feb 15 (the day everyone launched everything) was Gemini's turn, announcing the first LLM with 1 million token context window. * In May 2024 at Google I/O, Gemini 1.5 Pro announced a 2m token context window In parallel, open source/academia had to fight its own battle to keep up with the industrial cutting edge. Nous Research famously turned a reddit comment into YaRN, extending Llama 2 models to 128k context. So when Llama 3 dropped, the community was ready, and just weeks later, we had Llama3 with 4M+ context! A year ago we didn’t really have an industry standard way of measuring context utilization either: it’s all well and good to technically make an LLM generate non-garbage text at 1m tokens, but can you prove that the LLM actually retrieves and attends to information inside that long context? Greg Kamradt popularized the Needle In A Haystack chart which is now a necessary (if insufficient) benchmark — and it turns out we’ve solved that too in open source: Today's guest, Mark Huang, is the co-founder of Gradient, where they are building a full stack AI platform to power enterprise workflows and automations. They are also the team behind the first Llama3's 1M+ and 4M+ context window finetunes. Long Context Algorithms: RoPE, ALiBi, and Ring Attention Positional encodings allow the model to understand the relative position of tokens in the input sequence, present in what (upcoming guest!) Yi Tay affectionately calls the OG “Noam architecture”. But if we want to increase a model’s context length, these encodings need to gracefully extrapolate to longer sequences. ALiBi, used in models like MPT (see our "Context=Infinity" episode with the MPT leads, Jonathan Frankle and Abhinav), was one of the early approaches to this space. It lets the context window stretch as it grows, using a linearly decreasing penalty between attention weights of different positions; the further two tokens are, the higher the penalty. Of course, this isn’t going to work for usecases that actually require global attention across a long context. In more recent architectures and finetunes, RoPE (Rotary Position Embedding) encoding is more commonly used and is also what Llama3 was based on. RoPE uses a rotational matrix to encode positions, which empirically performs better for longer sequences. The main innovation from Gradient was to focus on tuning the theta hyperparameter that governs the frequency of the rotational encoding. Audio note: If you want the details, jump to 15:55 in the podcast (or scroll down to the transcript!) By carefully increasing theta as context length grew, they were able to scale Llama3 up to 1 million tokens and potentially beyond. Once you've scaled positional embeddings, there's still the issue of attention's quadratic complexity, and how longer and longer sequences impacts models speed and scaling abilities. Getting to 1-4M context window requires a fairly large amount of compute, so efficiency matters. Ring Attention was the other "one small trick that GPU clouds hate" that improves GPU utilization by allowing parallel computation and communication between GPUs. Gradient started from the EasyContext library as implementation of Ring Attention in PyTorch, since the original one was in JAX. Long Context Data: Curriculum Learning and Progressive Extension The use of curriculum learning when extending context was another new approach; rather than training Llama3 on the full 1 million token context from the start, they progressively increased the sequence length over the course of training. Intuitively, it allows the model to first learn to utilize shorter contexts before tackling the full length, but it only works if data gets more and more "tricky" in long context situation. For the generic pre-training corpus they used SlimPajama as a base, and concatenated texts to reach the target length, while monitoring for diversity in the data. Datasets that only required attending to the last few tokens, for instance, would fail to teach long-range reasoning. To fix that, they used synthetic data (another one of our Four Wars of AI!) with GPT-4 to augment their datasets by prompting it to expand on information or rephrase excerpts. Another paper we previously mentioned in this space is "Rephrasing The Web". Long Context Benchmarking: Beyond Needles Long context is cool, but does it work? Greg’s now-famous "needle in a haystack" (NIAH) test, which measures a model's ability to extract a piece of information embedded in a long context, is a clean standard that everyone uses to start, but it is a little simplistic and the community has since created many options to extend it: * RULER: Outside of various NIAH tests (single value, multiple values, etc) it also tests for things like "most frequent words" and "variable tracking", which is very helpful especially in coding use cases. * LooGLE: Focuses on three main area: scientific papers, Wikipedia articles, movie and TV scripts. "Timeline reorder" is an interesting challenge in their benchmark, which asks model to create a timeline out of events that happened out of order in the text. * Infinite Bench: First created in November 2023, most avg input tokens tasks are in the 100-200k tokens range across retrieval, Q&A, and code debugging. * ZeroSCROLLS: this comes with a public leaderboard where you can see models performance, as well as tasks that you can browse to get an idea. The 4M context size seemed to be the limit where things started to fall apart as far as performance goes, which is quite impressive! Show Notes * Mark Huang * Gradient * Chris Chang * HuggingFace Hub with Llama3 finetunes * Mad Men * Crusoe * Greg Kamradt's Needle in a Haystack * Chameleon paper * Charles Goddard (Mentioned in context with model merging) * Matei Zaharia * Phil Wang (lucidrains) * Wing Lian * Zhang Peiyuan * Yi * Scaling Laws of RoPE-based Extrapolation * ALiBi * YaRN * Ring Attention * Easy Context * StrongCompute * LoRa * RULER: What's the Real Context Size of Your Long-Context Language Models? * LooGLE: Can Long-Context Language Models Understand Long Contexts? * Infinite Bench * BAMBOO * ZeroSCROLLS: Zero-Shot CompaRison Over Long Language Sequences * DeepSeek paper * Multi-head Latent Attention Chapters * [00:00:01] Introductions * [00:01:28] Founding story of Gradient and its mission * [00:03:50] "Minimum viable agents" * [00:07:37] Differentiating ML and AI, focusing on out-of-domain generalization * [00:08:19] Extending Llama3 to 1M tokens * [00:11:41] Technical challenges with long context sequences * [00:14:30] Data quality and the importance of diverse datasets * [00:16:07] What's a theta value? * [00:18:27] RoPE vs Ring Attention vs ALiBi vs YaARN * [00:20:23] Why RingAttention matters * [00:22:47] How to refine datasets for context extension * [00:27:28] Multi-stage training data and avoiding overfitting to recent data * [00:28:10] The potential of using synthetic data in training * [00:31:21] Applying LoRa adapters to extend model capabilities * [00:34:45] Benchmarking long context models and evaluating their performance * [00:38:38] Pushing to 4M context and output quality degradation * [00:40:49] What do you need this context for? * [00:42:54] Impact of long context in chat vs Docs Summarization * [00:45:35] Future directions for long context models and multimodality * [00:48:01] How do you know what research matters? * [00:50:31] Routine for staying updated with AI research and industry news * [00:52:39] Deciding which AI developments to invest time in * [00:56:08] Request for collaboration and data set construction for long context Transcript Alessio [00:00:00]: Hey everyone, welcome to the Latent Space podcast. This is Alessio, partner and CTO-in-Residence at Decibel Partners, and I'm joined by my co-host Swyx, founder of Smol AI. Swyx [00:00:14]: Hey, and today we're in the remote studio with Mark Wang from Gradient. Welcome Mark. Mark [00:00:19]: Hey, glad to be here. It's really a great experience to be able to talk with you all. I know your podcast is really, really interesting and I always am listening to it every time you guys have a release. Alessio [00:00:31]: He's not a paid actor. He said that out of his own will. Swyx [00:00:34]: We'll give you the check later. So you're unusual in the sense that you and I go back to college. I don't exactly remember where we overlapped, but you know, we both went to Wharton. We went into the sort of quantitative developer realm. Mark [00:00:46]: Yeah, exactly. Kind of crazy, right? So it all goes full circle. I was a quant for quite a few years and then made it out into Silicon Valley and now we intersect again when it kind of feels like more or less the same, right? Like the AI wars, the trading wars back in the day too, to a certain extent and the gr





