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

1from typing import Any 

2 

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) 

17 

18 

19class NanogptArchitectureAdapter(ArchitectureAdapter): 

20 """Architecture adapter for NanoGPT models.""" 

21 

22 def __init__(self, cfg: Any) -> None: 

23 """Initialize the NanoGPT architecture adapter. 

24 

25 Args: 

26 cfg: The configuration object. 

27 """ 

28 super().__init__(cfg) 

29 

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 } 

68 

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 }