Skip to content

Commit 96c73a0

Browse files
committed
Address TransformerBridge traversal review feedback
1 parent 99e690e commit 96c73a0

5 files changed

Lines changed: 136 additions & 50 deletions

File tree

tests/integration/model_bridge/test_parent_module_traversal.py

Lines changed: 104 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -71,6 +71,35 @@ def _vit_config() -> ViTConfig:
7171
)
7272

7373

74+
def _bart_config() -> BartConfig:
75+
return BartConfig(
76+
vocab_size=32,
77+
d_model=16,
78+
encoder_layers=1,
79+
decoder_layers=1,
80+
encoder_attention_heads=4,
81+
decoder_attention_heads=4,
82+
encoder_ffn_dim=32,
83+
decoder_ffn_dim=32,
84+
max_position_embeddings=16,
85+
)
86+
87+
88+
def _hubert_config() -> HubertConfig:
89+
return HubertConfig(
90+
vocab_size=32,
91+
hidden_size=16,
92+
num_hidden_layers=1,
93+
num_attention_heads=4,
94+
intermediate_size=32,
95+
conv_dim=(8,),
96+
conv_stride=(2,),
97+
conv_kernel=(3,),
98+
num_conv_pos_embeddings=4,
99+
num_conv_pos_embedding_groups=2,
100+
)
101+
102+
74103
ARCHITECTURE_CASES = (
75104
ArchitectureCase(
76105
"gpt2-joint-qkv",
@@ -163,17 +192,7 @@ def _vit_config() -> ViTConfig:
163192
ArchitectureCase(
164193
"bart-encoder-decoder",
165194
BartForConditionalGeneration,
166-
lambda: BartConfig(
167-
vocab_size=32,
168-
d_model=16,
169-
encoder_layers=1,
170-
decoder_layers=1,
171-
encoder_attention_heads=4,
172-
decoder_attention_heads=4,
173-
encoder_ffn_dim=32,
174-
decoder_ffn_dim=32,
175-
max_position_embeddings=16,
176-
),
195+
_bart_config,
177196
"BartForConditionalGeneration",
178197
),
179198
ArchitectureCase(
@@ -202,18 +221,7 @@ def _vit_config() -> ViTConfig:
202221
ArchitectureCase(
203222
"hubert-audio",
204223
HubertForCTC,
205-
lambda: HubertConfig(
206-
vocab_size=32,
207-
hidden_size=16,
208-
num_hidden_layers=1,
209-
num_attention_heads=4,
210-
intermediate_size=32,
211-
conv_dim=(8,),
212-
conv_stride=(2,),
213-
conv_kernel=(3,),
214-
num_conv_pos_embeddings=4,
215-
num_conv_pos_embedding_groups=2,
216-
),
224+
_hubert_config,
217225
"HubertForCTC",
218226
),
219227
ArchitectureCase(
@@ -234,6 +242,8 @@ def _vit_config() -> ViTConfig:
234242
),
235243
)
236244

245+
ARCHITECTURE_CASE_BY_NAME = {case.name: case for case in ARCHITECTURE_CASES}
246+
237247

238248
def _named_identities(named_values: Any) -> dict[int, str]:
239249
return {id(value): name for name, value in named_values}
@@ -274,17 +284,7 @@ def test_parent_and_direct_traversal_have_identical_state(case: ArchitectureCase
274284

275285

276286
def test_parent_dtype_conversion_updates_container_owned_state() -> None:
277-
bart_config = BartConfig(
278-
vocab_size=32,
279-
d_model=16,
280-
encoder_layers=1,
281-
decoder_layers=1,
282-
encoder_attention_heads=4,
283-
decoder_attention_heads=4,
284-
encoder_ffn_dim=32,
285-
decoder_ffn_dim=32,
286-
max_position_embeddings=16,
287-
)
287+
bart_config = _bart_config()
288288
bridge = build_bridge_from_module(
289289
BartForConditionalGeneration(bart_config),
290290
"BartForConditionalGeneration",
@@ -305,17 +305,7 @@ def test_parent_dtype_conversion_updates_container_owned_state() -> None:
305305

306306

307307
def test_parent_assign_load_updates_container_owned_state() -> None:
308-
bart_config = BartConfig(
309-
vocab_size=32,
310-
d_model=16,
311-
encoder_layers=1,
312-
decoder_layers=1,
313-
encoder_attention_heads=4,
314-
decoder_attention_heads=4,
315-
encoder_ffn_dim=32,
316-
decoder_ffn_dim=32,
317-
max_position_embeddings=16,
318-
)
308+
bart_config = _bart_config()
319309
bridge = build_bridge_from_module(
320310
BartForConditionalGeneration(bart_config),
321311
"BartForConditionalGeneration",
@@ -336,3 +326,72 @@ def test_parent_assign_load_updates_container_owned_state() -> None:
336326
assert id(bridge.original_model.final_logits_bias) in {
337327
id(buffer) for buffer in parent.buffers()
338328
}
329+
330+
331+
@pytest.mark.parametrize(
332+
("case_name", "container_path", "state_name", "state_key"),
333+
(
334+
("bart-encoder-decoder", "", "final_logits_bias", "final_logits_bias"),
335+
("hubert-audio", "hubert", "masked_spec_embed", "hubert.masked_spec_embed"),
336+
),
337+
)
338+
def test_direct_assign_load_stays_current_after_apply(
339+
case_name: str, container_path: str, state_name: str, state_key: str
340+
) -> None:
341+
case = ARCHITECTURE_CASE_BY_NAME[case_name]
342+
config = case.config_factory()
343+
bridge = build_bridge_from_module(
344+
case.model_type(config),
345+
case.architecture,
346+
hf_config=config,
347+
dtype=torch.float32,
348+
device="cpu",
349+
model_name=f"tiny-{case.name}-direct-assign",
350+
)
351+
original_container = (
352+
bridge.original_model.get_submodule(container_path)
353+
if container_path
354+
else bridge.original_model
355+
)
356+
owner_container = (
357+
bridge._container_state_owners.get_submodule(container_path)
358+
if container_path
359+
else bridge._container_state_owners
360+
)
361+
replacement = torch.full_like(getattr(original_container, state_name), 7)
362+
363+
bridge.load_state_dict({state_key: replacement}, strict=False, assign=True)
364+
365+
assert getattr(owner_container, state_name) is getattr(original_container, state_name)
366+
bridge.cpu()
367+
assert torch.equal(getattr(original_container, state_name), replacement)
368+
369+
370+
@pytest.mark.parametrize(
371+
("case_name", "key_fragment", "expected_keys"),
372+
(
373+
("bert-nsp", "pooler", {"pooler.weight", "pooler.bias"}),
374+
("vit-bare-pooler", "pooler", {"pooler.weight", "pooler.bias"}),
375+
(
376+
"ast-audio-classifier",
377+
"classifier",
378+
{"classifier_ln.weight", "classifier_ln.bias"},
379+
),
380+
),
381+
)
382+
def test_task_head_state_dict_keys_are_not_reexpanded(
383+
case_name: str, key_fragment: str, expected_keys: set[str]
384+
) -> None:
385+
case = ARCHITECTURE_CASE_BY_NAME[case_name]
386+
config = case.config_factory()
387+
bridge = build_bridge_from_module(
388+
case.model_type(config),
389+
case.architecture,
390+
hf_config=config,
391+
dtype=torch.float32,
392+
device="cpu",
393+
model_name=f"tiny-{case.name}-state-dict-keys",
394+
)
395+
396+
actual_keys = {key for key in bridge.state_dict() if key_fragment in key}
397+
assert actual_keys == expected_keys

tests/unit/model_bridge/supported_architectures/test_vit_adapter.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -312,7 +312,7 @@ def test_bare_model_maps_pooler_without_root_name_collision(
312312
self, adapter: ViTArchitectureAdapter
313313
) -> None:
314314
adapter.prepare_model(self._bare_model_with_pooler())
315-
assert adapter.component_mapping["vision_pooler"].name == "pooler.dense"
315+
assert adapter.component_mapping["pooler"].name == "pooler.dense"
316316

317317
def test_bare_model_does_not_require_encoder_attribute(
318318
self, adapter: ViTArchitectureAdapter

transformer_lens/model_bridge/component_setup.py

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,15 @@ def _sync_original_container(self) -> None:
3232
original_container._parameters.update(self._parameters)
3333
original_container._buffers.update(self._buffers)
3434

35+
def _refresh_from_original_container(self) -> None:
36+
original_container = self.__dict__["_original_container"]
37+
for name in self._parameters:
38+
self._parameters[name] = original_container._parameters[name]
39+
for name in self._buffers:
40+
self._buffers[name] = original_container._buffers[name]
41+
3542
def _apply(self, fn: Any, recurse: bool = True) -> "_ContainerStateOwner":
43+
self._refresh_from_original_container()
3644
super()._apply(fn, recurse=recurse)
3745
self._sync_original_container()
3846
return self
@@ -59,6 +67,16 @@ def _load_from_state_dict(
5967
self._sync_original_container()
6068

6169

70+
def refresh_container_state_owners(bridge_module: nn.Module) -> None:
71+
"""Refresh registered container state from the original model tree."""
72+
root_owner = bridge_module._modules.get("_container_state_owners")
73+
if not isinstance(root_owner, _ContainerStateOwner):
74+
return
75+
for owner in root_owner.modules():
76+
if isinstance(owner, _ContainerStateOwner):
77+
owner._refresh_from_original_container()
78+
79+
6280
def replace_remote_component(
6381
replacement_component: nn.Module, remote_path: str, remote_model: RemoteModel
6482
) -> None:

transformer_lens/model_bridge/supported_architectures/vit.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -205,4 +205,4 @@ def prepare_model(self, hf_model: Any) -> None:
205205
prefix=prefix, with_classifier=with_classifier
206206
)
207207
if not with_classifier and getattr(hf_model, "pooler", None) is not None:
208-
self.component_mapping["vision_pooler"] = LinearBridge(name="pooler.dense")
208+
self.component_mapping["pooler"] = LinearBridge(name="pooler.dense")

transformer_lens/model_bridge/transformer_bridge.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,10 @@
3737
from transformer_lens.hook_points import HookIntrospectionMixin, HookPoint
3838
from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
3939
from transformer_lens.model_bridge.bridge_core import BridgeCore
40-
from transformer_lens.model_bridge.component_setup import set_original_components
40+
from transformer_lens.model_bridge.component_setup import (
41+
refresh_container_state_owners,
42+
set_original_components,
43+
)
4144
from transformer_lens.model_bridge.composition_scores import CompositionScores
4245
from transformer_lens.model_bridge.driver_protocol import (
4346
TensorLike,
@@ -404,6 +407,10 @@ def __getattr__(self, name: str) -> Any:
404407
# Use __dict__ directly to avoid recursion
405408
if "_modules" in self.__dict__ and name in self.__dict__["_modules"]: # type: ignore[arg-type]
406409
return self.__dict__["_modules"][name]
410+
adapter = self.__dict__.get("adapter")
411+
component_mapping = getattr(adapter, "component_mapping", None)
412+
if component_mapping is not None and name in component_mapping:
413+
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
407414
if "original_model" in self.__dict__ and self.__dict__["original_model"] is not None:
408415
try:
409416
name_split = name.split(".")
@@ -3791,8 +3798,8 @@ def _normalize_bridge_key_to_hf(self, key: str) -> str:
37913798
block_list_names = {"blocks", "L_blocks", "H_blocks", "encoder_blocks", "decoder_blocks"}
37923799
for tl_name, component in component_mapping.items():
37933800
if component.name and tl_name not in block_list_names:
3794-
# Skip if TL name is already a suffix of the HF path (avoids doubling).
3795-
if tl_name != component.name and not component.name.endswith("." + tl_name):
3801+
# Skip if TL name is already a segment of its HF path (avoids doubling).
3802+
if tl_name != component.name and tl_name not in component.name.split("."):
37963803
attr_to_hf[tl_name] = component.name
37973804

37983805
# Map block-level components (ln1, ln2, attn, mlp) for all block lists
@@ -3969,6 +3976,8 @@ def load_state_dict(self, state_dict, strict=True, assign=False):
39693976
)
39703977

39713978
result = self.original_model.load_state_dict(mapped_state_dict, strict=False, assign=assign)
3979+
if assign:
3980+
refresh_container_state_owners(self)
39723981
return type(result)(missing_keys=missing_keys, unexpected_keys=unexpected_keys)
39733982

39743983
def get_params(self):

0 commit comments

Comments
 (0)