Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions tests/integration/model_bridge/test_llada2_moe_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,15 +10,17 @@

from tests.integration.model_bridge.helpers import make_tiny_pair
from transformer_lens.model_bridge.generalized_components import MoEBridge
from transformer_lens.model_bridge.supported_architectures.dream import (
_register_default_rope_init,
from transformer_lens.model_bridge.supported_architectures._remote_code_compat import (
restore_default_rope_init,
)

MODEL_ID = "inclusionAI/LLaDA2.0-mini"


def _tiny_llada2_pair():
_register_default_rope_init()
restore_default_rope_init(
MODEL_ID, "modeling_llada2_moe.LLaDA2MoeModelLM", "LLaDA2MoeRotaryEmbedding"
)
from transformers import AutoConfig, AutoModelForCausalLM

cfg = AutoConfig.from_pretrained(MODEL_ID, trust_remote_code=True)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

Dream is Qwen2.5-shaped but bidirectional (diffusion): attention must stay
delegated to HF, generation phases are excluded, and the v5 rope shim must
restore the 'default' ROPE_INIT_FUNCTIONS entry the remote code looks up.
restore the 'default' rope init the remote code looks up.
"""
from typing import Any

Expand Down Expand Up @@ -79,12 +79,13 @@ def test_gated_mlp(self, adapter):


class TestDreamRopeShim:
def test_prepare_loading_registers_default_rope(self, adapter):
def test_prepare_loading_leaves_shared_rope_registry_alone(self, adapter):
"""The shim is module-local: a shared "default" entry would override every
native model's own default rope init on transformers>=5.17."""
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS

ROPE_INIT_FUNCTIONS.pop("default", None)
adapter.prepare_loading("Dream-org/Dream-v0-Instruct-7B", {})
assert "default" in ROPE_INIT_FUNCTIONS
assert "default" not in ROPE_INIT_FUNCTIONS

def test_v4_rope_matches_reference_formula(self):
class Cfg:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -67,12 +67,11 @@ def test_word_embeddings_path(self, adapter):
assert adapter.component_mapping["embed"].name == "model.word_embeddings"


def test_prepare_loading_registers_rope_shim(adapter=None):
def test_prepare_loading_leaves_shared_rope_registry_alone(adapter=None):
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS

ROPE_INIT_FUNCTIONS.pop("default", None)
LLaDA2MoeArchitectureAdapter(_make_cfg()).prepare_loading("inclusionAI/LLaDA2.0-mini", {})
assert "default" in ROPE_INIT_FUNCTIONS
assert "default" not in ROPE_INIT_FUNCTIONS


def test_factory_registration():
Expand Down
92 changes: 90 additions & 2 deletions tests/unit/model_bridge/test_remote_code_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
force_import_remote_class,
iter_remote_modeling_modules,
patch_init_weights_skip_loaded,
restore_default_rope_init,
retie_weights_keys_v5,
)

Expand Down Expand Up @@ -249,9 +250,96 @@ class Cfg:
assert inv_freq.shape == (8,)

def test_kwargs_only_path(self) -> None:
"""v4 also allowed configless calls with base/dim; dream registers the
helper globally, so arbitrary remote code may use that form."""
"""v4 also allowed configless calls with base/dim; remote code may use that form."""
inv_freq, scaling = compute_default_rope_inv_freq(base=10000.0, dim=8)
expected = 1.0 / (10000.0 ** (torch.arange(0, 8, 2, dtype=torch.int64).float() / 8))
torch.testing.assert_close(inv_freq, expected)
assert scaling == 1.0


class TestRestoreDefaultRopeInit:
MODULE = "transformers_modules.acme.zorbo.modeling_zorbo"

def _install_fake_remote_module(self, monkeypatch: pytest.MonkeyPatch) -> ModuleType:
"""A stand-in remote modeling module that imported the (default-less) v5 dict."""
import transformers.dynamic_module_utils as dmu
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS

class ZorboRotaryEmbedding:
pass

module = ModuleType(self.MODULE)
setattr(module, "ROPE_INIT_FUNCTIONS", ROPE_INIT_FUNCTIONS)
setattr(module, "ZorboRotaryEmbedding", ZorboRotaryEmbedding)
monkeypatch.setitem(sys.modules, self.MODULE, module)
monkeypatch.setattr(
dmu, "get_class_from_dynamic_module", lambda *args, **kwargs: ZorboRotaryEmbedding
)
return module

def test_patches_remote_module_only(self, monkeypatch: pytest.MonkeyPatch) -> None:
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS

module = self._install_fake_remote_module(monkeypatch)
restore_default_rope_init("acme/zorbo", "modeling_zorbo.ZorboModel", "ZorboRotaryEmbedding")

assert getattr(module, "ROPE_INIT_FUNCTIONS")["default"] is compute_default_rope_inv_freq
rotary_cls = getattr(module, "ZorboRotaryEmbedding")
assert rotary_cls.compute_default_rope_parameters is compute_default_rope_inv_freq
assert "default" not in ROPE_INIT_FUNCTIONS

def test_native_models_still_build_afterwards(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""transformers>=5.17 lets shared ROPE_INIT_FUNCTIONS entries override a
native model's own default rope init, so the shared dict must stay clean."""
from transformers import LlamaConfig, LlamaForCausalLM

self._install_fake_remote_module(monkeypatch)
restore_default_rope_init("acme/zorbo", "modeling_zorbo.ZorboModel", "ZorboRotaryEmbedding")

cfg = LlamaConfig(
hidden_size=16,
intermediate_size=32,
num_hidden_layers=1,
num_attention_heads=2,
vocab_size=32,
)
LlamaForCausalLM(cfg)

def test_noop_when_dynamic_module_unavailable(self, monkeypatch: pytest.MonkeyPatch) -> None:
import transformers.dynamic_module_utils as dmu

def boom(*args: Any, **kwargs: Any) -> None:
raise RuntimeError("offline / not a remote-code repo")

monkeypatch.setattr(dmu, "get_class_from_dynamic_module", boom)
restore_default_rope_init("acme/zorbo", "modeling_zorbo.ZorboModel", "ZorboRotaryEmbedding")

def test_patches_the_requested_revision(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""A pinned revision is its own module copy; only importing that revision
brings it into sys.modules, so the revision must be forwarded."""
import transformers.dynamic_module_utils as dmu
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS

default_copy = self._install_fake_remote_module(monkeypatch)
revision_name = "transformers_modules.acme.zorbo.abc123.modeling_zorbo"

def fake_import(ref: str, name: str, revision: str | None = None, **kwargs: Any) -> type:
if revision == "abc123":
module = ModuleType(revision_name)
setattr(module, "ROPE_INIT_FUNCTIONS", ROPE_INIT_FUNCTIONS)
monkeypatch.setitem(sys.modules, revision_name, module)
return type("ZorboModel", (), {})

monkeypatch.setattr(dmu, "get_class_from_dynamic_module", fake_import)
restore_default_rope_init(
"acme/zorbo", "modeling_zorbo.ZorboModel", "ZorboRotaryEmbedding", revision="abc123"
)

revision_copy = sys.modules[revision_name]
assert getattr(revision_copy, "ROPE_INIT_FUNCTIONS")["default"] is (
compute_default_rope_inv_freq
)
assert getattr(default_copy, "ROPE_INIT_FUNCTIONS")["default"] is (
compute_default_rope_inv_freq
)
assert "default" not in ROPE_INIT_FUNCTIONS
Original file line number Diff line number Diff line change
Expand Up @@ -138,9 +138,8 @@ def compute_default_rope_inv_freq(
Standard (unscaled) RoPE inverse frequencies with the v4 call contract the
remote code targets — ``(config, device) -> (inv_freq, attention_scaling)``
— plus v4's kwargs-only fallback (``base``/``dim`` passed directly).
Registration strategy stays with each adapter: dream re-registers the
global ``ROPE_INIT_FUNCTIONS["default"]``, ouro deliberately patches only
its own modeling module's copy.
Registered per remote modeling module, never in the shared
``ROPE_INIT_FUNCTIONS`` (see ``restore_default_rope_init``).
"""
if config is not None:
base = config.rope_theta
Expand All @@ -157,3 +156,35 @@ def compute_default_rope_inv_freq(
** (torch.arange(0, dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / dim)
)
return inv_freq, 1.0


def restore_default_rope_init(
model_name: str,
dotted_ref: str,
rotary_class_name: str,
revision: str | None = None,
) -> None:
"""Restore v4's ``ROPE_INIT_FUNCTIONS["default"]`` for one remote-code model.

``dotted_ref`` (``"<modeling_file>.<ClassName>"``) is force-imported at the
caller's ``revision`` (each revision is its own module copy); every
loaded copy of that modeling file gets its module-level
``ROPE_INIT_FUNCTIONS`` rebound to a copy with "default" restored, and its
``rotary_class_name`` gets ``compute_default_rope_parameters`` for v5's
``_init_weights``. The shared transformers dict must stay untouched: from
transformers 5.17, ``_init_weights`` lets its entries override every
native model's own ``compute_default_rope_parameters``.
"""
if force_import_remote_class(model_name, dotted_ref, revision=revision) is None:
return
for module in iter_remote_modeling_modules(dotted_ref.rsplit(".", 1)[0]):
rope_functions = getattr(module, "ROPE_INIT_FUNCTIONS", None)
if rope_functions is not None and "default" not in rope_functions:
setattr(
module,
"ROPE_INIT_FUNCTIONS",
{**rope_functions, "default": compute_default_rope_inv_freq},
)
rope_class = getattr(module, rotary_class_name, None)
if rope_class is not None and not hasattr(rope_class, "compute_default_rope_parameters"):
rope_class.compute_default_rope_parameters = staticmethod(compute_default_rope_inv_freq)
22 changes: 10 additions & 12 deletions transformer_lens/model_bridge/supported_architectures/dream.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,9 @@
wrapped projections; there is no reconstructed pattern hook.

The remote code targets transformers 4.46; v5 removed the "default" key
from ``ROPE_INIT_FUNCTIONS``, so ``prepare_loading`` re-registers it with
the v4 semantics (plain inverse-frequency rope, attention factor 1.0).
from ``ROPE_INIT_FUNCTIONS``, so ``prepare_loading`` restores it in the remote
modeling module with the v4 semantics (plain inverse-frequency rope,
attention factor 1.0).
"""

from typing import Any
Expand All @@ -25,8 +26,8 @@
LinearBridge,
)
from transformer_lens.model_bridge.supported_architectures._remote_code_compat import (
compute_default_rope_inv_freq,
force_import_remote_class,
restore_default_rope_init,
)
from transformer_lens.model_bridge.supported_architectures.qwen2 import (
Qwen2ArchitectureAdapter,
Expand Down Expand Up @@ -73,14 +74,6 @@ def from_model_config(cls: Any, model_config: Any) -> Any:
setattr(gen_cfg_cls, "_tl_from_model_config_patched", True)


def _register_default_rope_init() -> None:
"""Restore the global ``ROPE_INIT_FUNCTIONS["default"]`` entry v5 removed;
Dream's remote code (and llada2_moe's) looks it up by that key."""
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS

ROPE_INIT_FUNCTIONS.setdefault("default", compute_default_rope_inv_freq)


class DreamArchitectureAdapter(Qwen2ArchitectureAdapter):
"""Architecture adapter for DreamModel diffusion LMs."""

Expand Down Expand Up @@ -119,7 +112,12 @@ def _build_attention_bridge(self):

def prepare_loading(self, model_name: str, model_kwargs: dict) -> None:
"""Shim the remote code's two transformers-v4 dependencies."""
_register_default_rope_init()
restore_default_rope_init(
model_name,
"modeling_dream.DreamModel",
"DreamRotaryEmbedding",
revision=model_kwargs.get("revision"),
)
# DreamGenerationConfig.validate is a no-op with the v4 signature
# (is_init=False); v5 passes user_set_attributes. Replace with a
# kwargs-tolerant no-op.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,8 +35,8 @@
from transformer_lens.model_bridge.generalized_components.base import (
GeneralizedComponent,
)
from transformer_lens.model_bridge.supported_architectures.dream import (
_register_default_rope_init,
from transformer_lens.model_bridge.supported_architectures._remote_code_compat import (
restore_default_rope_init,
)


Expand Down Expand Up @@ -133,7 +133,12 @@ def __init__(self, cfg: Any) -> None:

def prepare_loading(self, model_name: str, model_kwargs: dict) -> None:
"""Restore the v4 'default' rope init the remote code looks up (Dream shim)."""
_register_default_rope_init()
restore_default_rope_init(
model_name,
"modeling_llada2_moe.LLaDA2MoeModelLM",
"LLaDA2MoeRotaryEmbedding",
revision=model_kwargs.get("revision"),
)
super().prepare_loading(model_name, model_kwargs)

def setup_hook_compatibility(self, bridge: Any) -> None:
Expand Down
Loading