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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
1"""Arcee 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 ArceeArchitectureAdapter(ArchitectureAdapter):
18 """Architecture adapter for Arcee models (ArceeForCausalLM / AFM-4.5B).
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.
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.
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``):
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
43 Weight processing handles these missing biases gracefully via
44 ProcessWeights._safe_get_tensor().
45 """
47 _testing_eager = None
49 def __init__(self, cfg: Any) -> None:
50 """Initialize the Arcee architecture adapter."""
51 super().__init__(cfg)
53 self._set_rms_rotary_defaults(gated=False)
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"
59 self.weight_processing_conversions = {
60 **self._qkvo_weight_conversions(),
61 }
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 }