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

11 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-21 19:27 +0000

1"""Arcee 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 ArceeArchitectureAdapter(ArchitectureAdapter): 

18 """Architecture adapter for Arcee models (ArceeForCausalLM / AFM-4.5B). 

19 

20 Arcee is a Llama-style dense decoder: pre-norm RMSNorm, rotary position 

21 embeddings (RoPE), grouped query attention (GQA), and no biases on any 

22 projection. The single distinguishing feature is the MLP: an *ungated* 

23 feed-forward block (``up_proj -> ReLU^2 -> down_proj``) using the squared-ReLU 

24 activation (HF ``hidden_act = "relu2"``) instead of the gated SiLU/GeLU used by 

25 Llama. The post-activation neurons are exposed via the MLP bridge's 

26 ``hook_post`` (``mlp.out.hook_in``), which is useful for inspecting the sparse 

27 activation structure ReLU^2 produces. 

28 

29 Structurally identical to Llama except for the ungated ReLU^2 MLP; unlike 

30 Apertus it uses standard ``input_layernorm`` / ``post_attention_layernorm`` 

31 names and has no Q/K normalization. 

32 

33 Optional Parameters (may not exist in state_dict): 

34 ------------------------------------------------- 

35 Arcee models do NOT have biases on attention or MLP projections 

36 (``attention_bias = false``, ``mlp_bias = false``): 

37 

38 - blocks.{i}.attn.b_Q / b_K / b_V / b_O - No bias on attention projections 

39 - blocks.{i}.mlp.b_in - No bias on MLP input (up_proj) 

40 - blocks.{i}.mlp.b_out - No bias on MLP output (down_proj) 

41 - blocks.{i}.ln1.b / ln2.b / ln_final.b - RMSNorm has no bias 

42 

43 Weight processing handles these missing biases gracefully via 

44 ProcessWeights._safe_get_tensor(). 

45 """ 

46 

47 _testing_eager = None 

48 

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

50 """Initialize the Arcee architecture adapter.""" 

51 super().__init__(cfg) 

52 

53 self._set_rms_rotary_defaults(gated=False) 

54 

55 # Use eager attention so output_attentions works for hook_attn_scores / 

56 # hook_pattern; SDPA does not support output_attentions. 

57 self.cfg.attn_implementation = "eager" 

58 

59 self.weight_processing_conversions = { 

60 **self._qkvo_weight_conversions(), 

61 } 

62 

63 self.component_mapping = { 

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

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

66 "blocks": BlockBridge( 

67 name="model.layers", 

68 submodules={ 

69 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg), 

70 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), 

71 "attn": PositionEmbeddingsAttentionBridge( 

72 name="self_attn", 

73 config=self.cfg, 

74 submodules={ 

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

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

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

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

79 }, 

80 requires_attention_mask=True, 

81 requires_position_embeddings=True, 

82 ), 

83 "mlp": self._ungated_mlp(), 

84 }, 

85 ), 

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

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

88 }