Coverage for transformer_lens/model_bridge/supported_architectures/nanogpt.py: 100%
10 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1from typing import Any
3from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion
4from transformer_lens.conversion_utils.param_processing_conversion import (
5 ParamProcessingConversion,
6)
7from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
8from transformer_lens.model_bridge.generalized_components import (
9 AttentionBridge,
10 BlockBridge,
11 EmbeddingBridge,
12 MLPBridge,
13 NormalizationBridge,
14 PosEmbedBridge,
15 UnembeddingBridge,
16)
19class NanogptArchitectureAdapter(ArchitectureAdapter):
20 """Architecture adapter for NanoGPT models."""
22 def __init__(self, cfg: Any) -> None:
23 """Initialize the NanoGPT architecture adapter.
25 Args:
26 cfg: The configuration object.
27 """
28 super().__init__(cfg)
30 self.weight_processing_conversions = {
31 "blocks.{i}.attn.q": ParamProcessingConversion(
32 tensor_conversion=RearrangeTensorConversion(
33 "d_model (3 n_head d_head) -> 3 n_head d_head d_model"
34 ),
35 source_key="transformer.h.{i}.attn.c_attn.weight",
36 ),
37 "blocks.{i}.attn.k": ParamProcessingConversion(
38 tensor_conversion=RearrangeTensorConversion(
39 "d_model (3 n_head d_head) -> 3 n_head d_head d_model"
40 ),
41 source_key="transformer.h.{i}.attn.c_attn.weight",
42 ),
43 "blocks.{i}.attn.v": ParamProcessingConversion(
44 tensor_conversion=RearrangeTensorConversion(
45 "d_model (3 n_head d_head) -> 3 n_head d_head d_model"
46 ),
47 source_key="transformer.h.{i}.attn.c_attn.weight",
48 ),
49 "blocks.{i}.attn.b_Q": ParamProcessingConversion(
50 tensor_conversion=RearrangeTensorConversion("(3 n_head d_head) -> 3 n_head d_head"),
51 source_key="transformer.h.{i}.attn.c_attn.bias",
52 ),
53 "blocks.{i}.attn.b_K": ParamProcessingConversion(
54 tensor_conversion=RearrangeTensorConversion("(3 n_head d_head) -> 3 n_head d_head"),
55 source_key="transformer.h.{i}.attn.c_attn.bias",
56 ),
57 "blocks.{i}.attn.b_V": ParamProcessingConversion(
58 tensor_conversion=RearrangeTensorConversion("(3 n_head d_head) -> 3 n_head d_head"),
59 source_key="transformer.h.{i}.attn.c_attn.bias",
60 ),
61 "blocks.{i}.attn.o": ParamProcessingConversion(
62 tensor_conversion=RearrangeTensorConversion(
63 "d_model (n_head d_head) -> n_head d_head d_model"
64 ),
65 source_key="transformer.h.{i}.attn.c_proj.weight",
66 ),
67 }
69 # Set up component mapping
70 self.component_mapping = {
71 "embed": EmbeddingBridge(name="transformer.wte"), # Word token embeddings
72 "pos_embed": PosEmbedBridge(name="transformer.wpe"), # Positional embeddings
73 "blocks": BlockBridge(
74 name="transformer.h", # Base path for blocks
75 submodules={
76 "ln1": NormalizationBridge(
77 name="ln_1", config=self.cfg
78 ), # Pre-attention layer norm
79 "ln2": NormalizationBridge(name="ln_2", config=self.cfg), # Pre-MLP layer norm
80 "attn": AttentionBridge(name="attn", config=self.cfg),
81 "mlp": MLPBridge(name="mlp"),
82 },
83 ),
84 "ln_final": NormalizationBridge(
85 name="transformer.ln_f", config=self.cfg
86 ), # Final layer norm
87 "unembed": UnembeddingBridge(name="lm_head"),
88 }