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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
1"""HunYuanDenseV1 architecture adapter."""
3from typing import Any
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)
17class HunYuanDenseV1ArchitectureAdapter(ArchitectureAdapter):
18 """Architecture adapter for HunYuanDenseV1 models."""
20 def __init__(self, cfg: Any) -> None:
21 super().__init__(cfg)
23 self._set_rms_rotary_defaults()
25 self.cfg.attn_implementation = "eager"
27 self.weight_processing_conversions = {
28 **self._qkvo_weight_conversions(),
29 }
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 }