← Back to Table of Contents

Chapter 25 β€” The transformers Library Internals

β€œUnderstanding how transformers (the library) is structured lets you read any model’s source, add custom architectures, and debug generation issues.”

Class Hierarchy

transformers Class Architecture
PreTrainedModel β€” Base class: from_pretrained(), save_pretrained(), generate(), gradient_checkpointing
LlamaModel (or GPT2Model, MistralModel, ...) β€” Core transformer: embeddings + N decoder layers + final norm
LlamaForCausalLM β€” Adds lm_head (linear projection to vocab). Computes cross-entropy loss.
GenerationMixin β€” generate() method: sampling, beam search, contrastive, speculative decode

Every model in the library follows this pattern. The Auto* classes simply look up the correct class from the model config:

1
2
3
4
5
# What AutoModelForCausalLM.from_pretrained() does internally:
# 1. Download config.json β†’ read "model_type": "llama"
# 2. Look up MODEL_FOR_CAUSAL_LM_MAPPING: "llama" β†’ LlamaForCausalLM
# 3. Instantiate LlamaForCausalLM(config)
# 4. Load safetensors weights into the model

Model Anatomy: LlamaForCausalLM

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
LlamaForCausalLM(
  (model): LlamaModel(
    (embed_tokens): Embedding(128256, 4096)           # token embeddings
    (layers): ModuleList(
      (0-31): 32 x LlamaDecoderLayer(
        (self_attn): LlamaSdpaAttention(
          (q_proj): Linear(4096, 4096, bias=False)     # Q projection
          (k_proj): Linear(4096, 1024, bias=False)     # K projection (GQA: fewer heads)
          (v_proj): Linear(4096, 1024, bias=False)     # V projection (GQA)
          (o_proj): Linear(4096, 4096, bias=False)     # output projection
          (rotary_emb): LlamaRotaryEmbedding()         # RoPE
        )
        (mlp): LlamaMLP(
          (gate_proj): Linear(4096, 14336, bias=False) # SwiGLU gate
          (up_proj): Linear(4096, 14336, bias=False)   # SwiGLU up
          (down_proj): Linear(14336, 4096, bias=False) # SwiGLU down
        )
        (input_layernorm): LlamaRMSNorm(4096)          # pre-attention norm
        (post_attention_layernorm): LlamaRMSNorm(4096) # pre-FFN norm
      )
    )
    (norm): LlamaRMSNorm(4096)                         # final norm
  )
  (lm_head): Linear(4096, 128256, bias=False)          # maps to vocab logits
)

Forward Pass Walkthrough

Input token IDs: B Γ— T
After embed_tokens: B Γ— T Γ— 4096
Γ— 32 LlamaDecoderLayers: B Γ— T Γ— 4096
After final norm: B Γ— T Γ— 4096
After lm_head: B Γ— T Γ— 128256
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
# Simplified forward of LlamaForCausalLM
def forward(self, input_ids, attention_mask=None, labels=None, past_key_values=None):
    # 1. Embeddings
    hidden_states = self.model.embed_tokens(input_ids)  # [B, T] β†’ [B, T, d]
    
    # 2. Decoder layers
    for layer in self.model.layers:
        hidden_states = layer(
            hidden_states,
            attention_mask=attention_mask,
            past_key_values=past_key_values,
        )
    
    # 3. Final norm + lm_head
    hidden_states = self.model.norm(hidden_states)       # [B, T, d]
    logits = self.lm_head(hidden_states)                 # [B, T, vocab_size]
    
    # 4. Loss (if labels provided)
    loss = None
    if labels is not None:
        # Shift: predict next token
        shift_logits = logits[..., :-1, :].contiguous()
        shift_labels = labels[..., 1:].contiguous()
        loss = F.cross_entropy(
            shift_logits.view(-1, self.config.vocab_size),
            shift_labels.view(-1),
        )
    
    return CausalLMOutputWithPast(loss=loss, logits=logits,
                                   past_key_values=past_key_values)

KV-Cache in transformers

During generation, the model caches key/value tensors to avoid recomputation:

1
2
3
4
5
6
7
8
9
10
11
12
13
# First call (prefill): process full prompt
outputs = model(input_ids, use_cache=True)
past_key_values = outputs.past_key_values  # DynamicCache object

# Subsequent calls (decode): process one new token at a time
for _ in range(max_new_tokens):
    outputs = model(
        next_token_id.unsqueeze(0),         # [1, 1] β€” single token
        past_key_values=past_key_values,    # reuse cached K, V
        use_cache=True,
    )
    past_key_values = outputs.past_key_values  # updated cache
    next_token_id = outputs.logits[:, -1, :].argmax(dim=-1)

The past_key_values is a DynamicCache containing per-layer K and V tensors:

1
2
3
4
5
# Inspect cache structure
cache = outputs.past_key_values
print(len(cache))              # 32 (one per layer)
print(cache[0][0].shape)       # K: [B, n_kv_heads, T_cached, head_dim]
print(cache[0][1].shape)       # V: [B, n_kv_heads, T_cached, head_dim]

Extending transformers: Custom Models

Register a custom architecture:

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
from transformers import PreTrainedModel, PretrainedConfig

class MyConfig(PretrainedConfig):
    model_type = "my_transformer"
    def __init__(self, d_model=1024, n_layers=12, n_heads=16, vocab_size=32000, **kwargs):
        self.d_model = d_model
        self.n_layers = n_layers
        self.n_heads = n_heads
        super().__init__(vocab_size=vocab_size, **kwargs)

class MyModel(PreTrainedModel):
    config_class = MyConfig
    
    def __init__(self, config):
        super().__init__(config)
        self.embed = nn.Embedding(config.vocab_size, config.d_model)
        self.layers = nn.ModuleList([
            TransformerBlock(config) for _ in range(config.n_layers)
        ])
        self.norm = nn.RMSNorm(config.d_model)
        self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
    
    def forward(self, input_ids, **kwargs):
        x = self.embed(input_ids)
        for layer in self.layers:
            x = layer(x)
        return self.lm_head(self.norm(x))

# Register and use
MyConfig.register_for_auto_class()
MyModel.register_for_auto_class("AutoModelForCausalLM")

# Now AutoModelForCausalLM.from_pretrained works with your model

Attention Implementations

transformers supports multiple attention backends, selectable via attn_implementation:

Backend Key Speed Memory Requirements
"eager" Manual PyTorch Baseline O(TΒ²) Always available
"sdpa" F.scaled_dot_product_attention 2Γ— Depends on backend PyTorch β‰₯ 2.0
"flash_attention_2" Flash Attention 2 2–4Γ— O(T) flash-attn package
1
2
3
4
5
model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.1-8B-Instruct",
    attn_implementation="flash_attention_2",  # or "sdpa", "eager"
    torch_dtype=torch.bfloat16,
)

What’s Next

How do we know if a model is actually good? The next chapter covers evaluation and benchmarks β€” the metrics, datasets, and leaderboards used to measure LLM quality.

← Previous: Chapter 28 β€” The Hugging Face Ecosystem Β· Next: Chapter 30 β€” Evaluation & Benchmarks β†’


Last updated: April 2026