Most open language models are pretrained at a few thousand tokens and then extended to tens or hundreds of thousands in a later, much cheaper phase. The mathematics of why a rotary position embedding breaks beyond its training length, and how Position Interpolation, NTK-aware scaling and YaRN repair it, is derived step by step in the context extension mathematics. This article is about everything around that formula: the project that turns an 8K model into a 64K model that actually uses its context.

That project has four parts, and skipping any of them produces a model that accepts long inputs but does not understand them. You choose and configure a position scheme, build a data mixture with genuinely long documents, run staged continued pretraining on infrastructure that can hold long sequences, and pass evaluation gates that measure use of distant tokens rather than the absence of errors. Then you pay for the result at serving time in KV-cache memory. Each part is explained below with code, a worked example and the failures that show up in real runs.

Advertisement

What has to be true for long context to work

A model can use a 64K context only when three independent conditions hold. First, the positions it sees must be in distribution: every rotary frequency must produce angles the attention layers were trained on, which is what the scaling method arranges. Second, the model must have learned to put attention weight on far-away tokens when they matter, which only training on long documents teaches. Third, the system must afford the memory and compute of long sequences in both training and serving.

Training-free scaling fixes the first condition only. That is why a model with a scaling factor bolted on often reads fluent text at long lengths yet fails to retrieve a fact from the middle. The diagram shows the whole pipeline; the gates at the bottom left are what decide whether a checkpoint moves forward.

A context-extension project: four gates between a short model and a long oneBase modeltrained at 8KPosition configYaRN, PI, theta, llama3Long-data mixtureupsampled long docsContinued pretraining16K, 32K, then 64KContext parallelring or all-to-allEval gatesposition NLL, RULER, shortLong SFTkeep the same configServing limitsKV memory, max lengthProductionmeasured effective lengthshards seqcheckpointpassfail: more dataor smaller factorChanging the position encoding is a config edit; the context only becomes usable after training and gates.
Configuration makes long inputs legal; data, training and gates make them useful; serving limits decide what you can afford to offer.

Step one: choose and configure the position scheme

For RoPE models, the practical options are those that Hugging Face Transformers implements as named rope types, because then the same config works in training, evaluation and serving. The trade-offs, in engineering terms:

Methodrope_typeNeeds trainingShort-context costNotes
Position InterpolationlinearYes, short fine-tuneSome, all frequencies squeezedSimple; weaker at large factors
Dynamic NTKdynamicNoNone below the original lengthInference-time stopgap; scale grows with input
YaRNyarnRecommendedSmallInterpolates by frequency band and adds attention temperature
Llama 3.1 stylellama3YesSmallBand-wise scaling with low and high frequency factors
LongRoPElongropeYesUses separate short factorsPer-dimension factors found by search
Base increase (ABF)default, larger rope_thetaYes, moreSmall after trainingChanges the model itself; used in many long pretraining runs

The configuration key changed between library generations: Transformers 4.x used rope_scaling, and 5.x uses rope_parameters, which also carries rope_theta. For YaRN, factor is the ratio of target to original length, original_max_position_embeddings is the pretraining length, and beta_fast and beta_slow default to 32 and 1. The exact frequency formulas are in YaRN and NTK scaling.

import torch
from transformers import AutoConfig, AutoModelForCausalLM

base = "your-org/base-8k"          # a RoPE model pretrained at 8,192 tokens
target = 65_536
yarn = {"rope_type": "yarn", "factor": target / 8192,
        "original_max_position_embeddings": 8192}

cfg = AutoConfig.from_pretrained(base)
if hasattr(cfg, "rope_parameters"):             # transformers 5.x
    cfg.rope_parameters = {**(cfg.rope_parameters or {}), **yarn}   # keeps rope_theta
else:                                           # 4.x uses rope_scaling
    cfg.rope_scaling = yarn
cfg.max_position_embeddings = target

model = AutoModelForCausalLM.from_pretrained(
    base, config=cfg, dtype=torch.bfloat16,
    attn_implementation="flash_attention_2")
model.gradient_checkpointing_enable()

One subtle trap: static scaling applies at every length, including short prompts. Qwen's model cards, for example, note that vLLM supports only static YaRN, which may reduce quality on shorter texts, and advise adding the scaling config only when long inputs are needed. Measure short-context quality with and without the config before shipping it.

Advertisement

Step two: build the long-data mixture

Long documents are rare. In a typical web crawl most documents are a few thousand tokens or less, and the ones above 32K concentrate in books, code repositories, scientific papers and legal or financial filings. The published recipes for continued pretraining converge on the same shape: upsample long documents within each domain rather than switching to a book-only diet, keep the overall domain mixture close to the pretraining mix, and keep a share of short data so short-context skills do not regress. The budget is modest: on the order of billions of tokens, not trillions.

  • Length-aware upsampling. Bucket documents by token length per domain and sample long buckets more often, while holding the domain weights fixed; changing both at once confounds your evaluation.
  • Real dependency, not concatenation. Packing unrelated short documents into a 64K sequence teaches nothing about distance. Prefer naturally long sources, whole repositories ordered by imports, and books; add synthetic long tasks such as multi-document question answering only as a minority.
  • Document masking. When you do pack, restart position ids per document and use a variable-length attention kernel so tokens never attend across document boundaries.
  • Deduplication at length. Long-document sources carry near-duplicates, such as book editions and forked repositories, which the upsampling then multiplies; deduplicate before you upsample.

Step three: staged continued pretraining

Attention memory with FlashAttention grows linearly with sequence length, but activations for the whole layer stack still scale with tokens per sequence, and compute per token rises because each token attends to more predecessors. A 64K sequence on a 7B to 8B model generally does not fit on one GPU without help. The standard toolkit is FlashAttention, activation checkpointing, sharded optimizer state, and context parallelism, which splits one sequence across GPUs and exchanges keys and values either around a ring or with all-to-all collectives. How the ring variant overlaps communication with compute is covered in ring attention on GPUs.

Stage the length rather than jumping straight to the target. Short stages are cheaper per token, adapt the model to the new position config, and give you a checkpoint to fall back to. Keep the learning rate well below the pretraining peak, with a short warmup, because the goal is adaptation, not new knowledge.

from transformers import DataCollatorWithFlattening

# stage schedule: (sequence length, tokens) -- lengths grow, lr stays low
STAGES = [(16_384, 1.0e9), (32_768, 1.0e9), (65_536, 1.5e9)]
collate = DataCollatorWithFlattening()   # input_ids, labels, position_ids per document

for seq_len, budget in STAGES:
    loader = packed_loader(mixture, seq_len, collate)   # documents packed to seq_len
    seen = 0
    for batch in loader:
        # position_ids restart at 0 for each packed document, so a
        # varlen attention kernel never attends across documents
        loss = model(**to_device(batch)).loss
        loss.backward()
        clip_grad_norm_(model.parameters(), 1.0)
        opt.step(); sched.step(); opt.zero_grad(set_to_none=True)
        seen += batch["input_ids"].numel()
        if seen >= budget:
            break
    save_checkpoint(model, f"ctx{seq_len}")
    run_eval_gates(model, seq_len)      # stop here if a gate fails

DataCollatorWithFlattening concatenates a batch into one row, marks boundaries in the labels with -100 and returns position ids that restart at zero for each document; it can also return the cumulative sequence lengths that FlashAttention's variable-length kernel uses. Check that your model's attention implementation honours these, or packing silently becomes cross-document attention.

Step four: evaluation gates

Long-context evaluation fails in a characteristic way: easy tests pass while the ability you care about is missing. Use gates of increasing difficulty and do not promote a checkpoint until all pass.

  1. Loss by position. On held-out long documents, average next-token loss per position bucket should keep falling or stay flat as position grows. A rise past the original length means positions are still out of distribution. It is necessary, not sufficient: loss can be good while retrieval is poor.
  2. Retrieval. Needle-in-a-haystack tests at many depths and lengths. Passing them is also not sufficient; single-needle retrieval is the easiest long-context task.
  3. Harder synthetic suites. RULER-style tasks with multiple keys, variable tracking and aggregation. The length where accuracy falls below your threshold is the effective context, and it is often well below the advertised one.
  4. Short-context regression. Your usual benchmarks at normal lengths, compared with the base model. Extension that costs a few points here is common and must be a deliberate trade.
  5. Real tasks. Long documents from your own domain with questions that need information from several places.
import torch
import torch.nn.functional as F

@torch.no_grad()
def nll_by_position(model, ids, edges=(0, 4096, 8192, 16384, 32768, 65536), chunk=4096):
    """Mean next-token loss per position bucket for one long document (Llama-style model)."""
    ids = ids.cuda()
    hidden = model.model(ids[None]).last_hidden_state[0]      # [T, d], no giant logits
    losses = []
    for s in range(0, ids.numel() - 1, chunk):
        e = min(s + chunk, ids.numel() - 1)
        logits = model.lm_head(hidden[s:e]).float()
        losses.append(F.cross_entropy(logits, ids[s + 1:e + 1], reduction="none"))
    nll = torch.cat(losses)
    return {f"{a}-{b}": nll[a:b].mean().item()
            for a, b in zip(edges, edges[1:]) if a < nll.numel()}

The helper avoids materialising logits for the whole sequence, which at 64K tokens and a large vocabulary would need tens of gigabytes, by applying the output head chunk by chunk.

Serving cost: the KV cache bill

Every token in the context keeps its keys and values in the KV cache for every layer. Per token that is 2 times layers times KV heads times head dimension times bytes per value. Worked example with a Llama 3.1 8B shape, 32 layers, 8 KV heads through grouped-query attention and head dimension 128, in BF16: 2 x 32 x 8 x 128 x 2 bytes is 131,072 bytes, 128 KiB per token. A single 131,072-token sequence therefore holds 16 GiB of cache, as much as the model weights. Eight concurrent users at full length need 128 GiB before activations. The derivation and quantisation options are in the KV cache mathematics.

Prefill is the other cost: attention work grows with the square of the prompt, so time to first token for a 128K prompt is far higher than for 8K, even with FlashAttention. In practice, set the serving engine's maximum model length, for example vLLM's --max-model-len, to what you need rather than the maximum the model supports, since it bounds memory planning; use prefix caching for repeated long documents; and consider KV-cache quantisation when memory, not compute, limits concurrency.

Failure modes and trade-offs

SymptomLikely causeResponse
Fluent but misses facts in the middleScaling without enough long trainingMore long, genuinely dependent data; check the effective length
Loss rises after the original lengthWrong factor or original length in configRecompute factor; confirm the loaded config
Short benchmarks dropMixture too long-heavy or static scalingRestore short data share; test dynamic or no scaling for short prompts
Works in training, fails in servingEngine ignores or overrides the rope configCompare logits between training code and the engine on one prompt
Out of memory at 64KNo context parallelism or checkpointingEnable both; reduce micro-batch to one sequence
Cross-document leakagePacking without position resetsFlattening collator with a varlen kernel

Finally, question whether you need a longer window at all. Retrieval that puts the right 8K tokens in the prompt is often cheaper and more accurate than a 128K window, and summarise-then-answer pipelines handle very long inputs at bounded cost. Long context strategies compares these families. Extension wins when the task needs the model to connect details scattered across a long document that retrieval would split apart.

What to do next

  1. Write down the target length and measure the base model's loss by position at that length with no scaling, as a baseline.
  2. Pick one rope type, compute the factor from the original length, and verify the loaded config in both the training and serving stacks.
  3. Build a length-bucketed mixture that upsamples long documents per domain and keeps a short-data share.
  4. Enable FlashAttention, activation checkpointing and context parallelism, then run staged lengths with a low learning rate.
  5. Gate each stage on loss by position, retrieval, a harder synthetic suite and short-context regression.
  6. Size the KV cache for your concurrency and set the serving maximum length accordingly.
  7. Publish the measured effective context, not the configured one.
Key takeaway: Extending an LLM's context is a data and systems project with a small configuration change at its centre. The rope scaling config makes long positions legal, but only continued pretraining on genuinely long, well-mixed, document-masked data teaches the model to use them, and only staged training with context parallelism makes that affordable. Gate every checkpoint on position-wise loss, retrieval, harder multi-hop suites and short-context regression, report the effective length you measured, and budget the KV cache honestly, because at long lengths it rivals the weights in size.