Coverage for transformer_lens/model_bridge/supported_architectures/hunyuan_v1_dense.py: 100%

10 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""HunYuanDenseV1 architecture adapter.""" 

2 

3from typing import Any 

4 

5from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

6from transformer_lens.model_bridge.generalized_components import ( 

7 BlockBridge, 

8 EmbeddingBridge, 

9 LinearBridge, 

10 PositionEmbeddingsAttentionBridge, 

11 RMSNormalizationBridge, 

12 RotaryEmbeddingBridge, 

13 UnembeddingBridge, 

14) 

15 

16 

17class HunYuanDenseV1ArchitectureAdapter(ArchitectureAdapter): 

18 """Architecture adapter for HunYuanDenseV1 models.""" 

19 

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

21 """Initialize the HunYuanDenseV1 architecture adapter.""" 

22 super().__init__(cfg) 

23 

24 self._set_rms_rotary_defaults() 

25 

26 self.cfg.attn_implementation = "eager" 

27 

28 self.weight_processing_conversions = { 

29 **self._qkvo_weight_conversions(), 

30 } 

31 

32 self.component_mapping = { 

33 "embed": EmbeddingBridge(name="model.embed_tokens"), 

34 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"), 

35 "blocks": BlockBridge( 

36 name="model.layers", 

37 submodules={ 

38 "ln1": RMSNormalizationBridge( 

39 name="input_layernorm", 

40 config=self.cfg, 

41 ), 

42 "ln2": RMSNormalizationBridge( 

43 name="post_attention_layernorm", 

44 config=self.cfg, 

45 ), 

46 "attn": PositionEmbeddingsAttentionBridge( 

47 name="self_attn", 

48 config=self.cfg, 

49 submodules={ 

50 "q": LinearBridge(name="q_proj"), 

51 "k": LinearBridge(name="k_proj"), 

52 "v": LinearBridge(name="v_proj"), 

53 "o": LinearBridge(name="o_proj"), 

54 "q_norm": RMSNormalizationBridge( 

55 name="query_layernorm", config=self.cfg 

56 ), 

57 "k_norm": RMSNormalizationBridge(name="key_layernorm", config=self.cfg), 

58 }, 

59 requires_attention_mask=True, 

60 requires_position_embeddings=True, 

61 ), 

62 "mlp": self._gated_mlp(), 

63 }, 

64 ), 

65 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg), 

66 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg), 

67 }