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-21 19:27 +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 super().__init__(cfg) 

22 

23 self._set_rms_rotary_defaults() 

24 

25 self.cfg.attn_implementation = "eager" 

26 

27 self.weight_processing_conversions = { 

28 **self._qkvo_weight_conversions(), 

29 } 

30 

31 self.component_mapping = { 

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

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

34 "blocks": BlockBridge( 

35 name="model.layers", 

36 submodules={ 

37 "ln1": RMSNormalizationBridge( 

38 name="input_layernorm", 

39 config=self.cfg, 

40 ), 

41 "ln2": RMSNormalizationBridge( 

42 name="post_attention_layernorm", 

43 config=self.cfg, 

44 ), 

45 "attn": PositionEmbeddingsAttentionBridge( 

46 name="self_attn", 

47 config=self.cfg, 

48 submodules={ 

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

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

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

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

53 "q_norm": RMSNormalizationBridge( 

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

55 ), 

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

57 }, 

58 requires_attention_mask=True, 

59 requires_position_embeddings=True, 

60 ), 

61 "mlp": self._gated_mlp(), 

62 }, 

63 ), 

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

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

66 }