Chapter 12 β Mid-Training & Continued Pre-Training
βMid-training is the bridge between raw pre-training and task-specific fine-tuning β where models gain long-context ability, domain expertise, and code fluency.β
What Is Mid-Training?
Mid-training (also called βcontinued pre-trainingβ or βpost-training phase 0β) is the stage between base pre-training and instruction fine-tuning. It uses the same next-token prediction objective as pre-training, but with a curated data mix designed for specific capabilities.
Long-Context Extension
One of the most important mid-training applications: extending a modelβs context window from 8K to 128K+ tokens.
The Challenge
Models pre-trained with 8K context canβt simply be used with 128K inputs β the positional encoding (RoPE) hasnβt seen those positions, and attention patterns degrade. The solution: progressive context extension during mid-training.
RoPE Scaling Methods
| Method | Description | Models |
|---|---|---|
| Position Interpolation | Scale frequencies by L_new/L_old β compress positions to fit | CodeLLaMA (16Kβ100K) |
| NTK-aware Scaling | Adjust RoPE base frequency instead of scaling | YaRN |
| YaRN | NTK + attention scaling + temperature correction | Mistral, various |
| ABF (Adjusted Base Frequency) | Increase ΞΈ_base (e.g., 500K β 8M) | LLaMA-3.1 (128K) |
LLaMA-3.1 increased RoPE base frequency from 500,000 to 8,000,000 and trained on progressively longer sequences:
The attention matrix goes from [B, H, 8K, 8K] to [B, H, 128K, 128K] β thatβs 256Γ more computation. This is why Flash Attention (Ch 4) and efficient KV-cache strategies (Ch 17β18) are essential.
Domain Adaptation
Mid-training adapts a general model to a specific domain by continuing to train on domain-specific data while mixing in general data to prevent forgetting.
Key Examples
Data Mixing to Avoid Catastrophic Forgetting
The critical challenge: training on domain data without forgetting general capabilities. The solution is data mixing β blending domain data with a fraction of general pre-training data:
1
2
3
4
5
6
7
# Typical mid-training data mix
data_mix = {
"domain_code": 0.60, # 60% new domain data
"general_web": 0.25, # 25% general data (prevent forgetting)
"high_quality_curated": 0.10, # 10% Wikipedia, books, math
"instruction_adjacent": 0.05, # 5% natural QA-like patterns
}
| Mix Strategy | Domain Quality | General Retention | Risk |
|---|---|---|---|
| 100% domain | High | Low | Catastrophic forgetting |
| 80/20 domain/general | High | Medium | Mild degradation on general tasks |
| 60/40 domain/general | Good | Good | Best balance for most cases |
| Progressive (β domain over time) | High | Good | More complex to implement |
Annealing
Near the end of mid-training (or the end of pre-training), learning rate annealing with high-quality data can significantly improve model capabilities:
- Cosine decay continues to min_lr
- Final data mix unchanged
- Good but not optimal final quality
- Aggressive LR drop (β 0) in final 5β10% of training
- Switch to high-quality data only
- Significant benchmark improvements (~1β3%)
LLaMA-3 used annealing on the final 40M tokens with a curated high-quality data mix, which notably improved benchmarks.
Continued Pre-Training Setup
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
from datasets import load_dataset
# Load base model
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3.1-8B",
torch_dtype=torch.bfloat16,
attn_implementation="flash_attention_2",
)
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B")
# Load domain data (e.g., code)
dataset = load_dataset("bigcode/starcoderdata", split="train", streaming=True)
training_args = TrainingArguments(
output_dir="./llama-code-mid-training",
per_device_train_batch_size=2,
gradient_accumulation_steps=16,
learning_rate=2e-5, # Lower than pre-training peak (3e-4)
lr_scheduler_type="cosine",
warmup_steps=500,
max_steps=50000,
bf16=True,
gradient_checkpointing=True, # Save memory at cost of ~30% speed
logging_steps=10,
save_steps=1000,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=dataset,
)
trainer.train()
Key differences from pre-training:
- Lower learning rate (typically 10β100Γ lower than pre-training peak)
- Fewer steps (thousands, not millions)
- Curated data (quality over quantity)
- Often single-stage (no complex parallelism β fits on fewer GPUs)
Fill-in-the-Middle (FIM) Training
Code mid-training often adds a fill-in-the-middle objective alongside standard next-token prediction. This enables code infilling (autocomplete in the middle of a function):
1
2
3
4
5
6
7
# Standard causal: predict left-to-right
# "def add(a, b):\n return a + b"
# FIM transformation (50% of training samples):
# "<|fim_prefix|>def add(a, b):\n <|fim_suffix|>\n<|fim_middle|>return a + b"
# The model learns to generate the middle given prefix + suffix
This is how Copilot-style code completion works β the model sees code before and after the cursor, and fills in the gap.
Whatβs Next
After mid-training produces a capable base model, supervised fine-tuning teaches it to follow instructions and behave as a helpful assistant.
β Previous: Chapter 11 β Optimizers & Loss Functions Β· Next: Chapter 13 β Fine-Tuning & Adaptation β
Last updated: April 2026