Coverage for transformer_lens/model_bridge/supported_architectures/flex_olmo.py: 100%
6 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
1"""FlexOlmo architecture adapter.
3AllenAI's FlexOlmo (``FlexOlmoForCausalLM``, NeurIPS 2025): federated MoE
4built by merging independently trained OLMo-2 experts, enabling
5inference-time data opt-out by expert selection. Structurally it is the
6exact union of the two OLMo variants already supported: OLMo-2's post-norm
7blocks and full-width q/k norms, with OLMoE's batched-parameter sparse MoE
8(gate_up_proj/down_proj as 3D tensors behind a top-k router) in place of
9the dense MLP. The router is a raw-parameter module (not nn.Linear), so it
10is wrapped for hooks as a plain delegated component.
11"""
13from transformer_lens.model_bridge.generalized_components import MoEBridge
14from transformer_lens.model_bridge.generalized_components.base import (
15 GeneralizedComponent,
16)
17from transformer_lens.model_bridge.supported_architectures.olmo2 import (
18 Olmo2ArchitectureAdapter,
19)
22class FlexOlmoArchitectureAdapter(Olmo2ArchitectureAdapter):
23 """Architecture adapter for FlexOlmoForCausalLM models."""
25 def _build_mlp_bridge(self) -> MoEBridge:
26 """Batched-expert sparse MoE, delegated to HF's native forward."""
27 return MoEBridge(
28 name="mlp",
29 config=self.cfg,
30 submodules={
31 "gate": GeneralizedComponent(name="gate"),
32 },
33 )