[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
This commit is contained in:
parent
6ce495a42d
commit
d2d758d0d8
2 changed files with 61 additions and 57 deletions
|
|
@ -759,9 +759,7 @@ def _rebind_in_already_imported_modules(
|
|||
setattr(mod, attr_name, new_obj)
|
||||
count += 1
|
||||
except Exception as exc:
|
||||
logger.debug(
|
||||
"Could not rebind %s in %s: %s", attr_name, mod_name, exc
|
||||
)
|
||||
logger.debug("Could not rebind %s in %s: %s", attr_name, mod_name, exc)
|
||||
return count
|
||||
|
||||
|
||||
|
|
@ -856,14 +854,14 @@ def _install_fast_path_hooks(event_queue: Any) -> None:
|
|||
# Reuse the existing wheel-first installer. It does its own
|
||||
# idempotency check via `__import__("causal_conv1d")`.
|
||||
_install_package_wheel_first(
|
||||
event_queue=eq,
|
||||
import_name="causal_conv1d",
|
||||
display_name="causal-conv1d",
|
||||
pypi_name="causal-conv1d",
|
||||
pypi_version=_CAUSAL_CONV1D_PACKAGE_VERSION,
|
||||
filename_prefix="causal_conv1d",
|
||||
release_tag=_CAUSAL_CONV1D_RELEASE_TAG,
|
||||
release_base_url=(
|
||||
event_queue = eq,
|
||||
import_name = "causal_conv1d",
|
||||
display_name = "causal-conv1d",
|
||||
pypi_name = "causal-conv1d",
|
||||
pypi_version = _CAUSAL_CONV1D_PACKAGE_VERSION,
|
||||
filename_prefix = "causal_conv1d",
|
||||
release_tag = _CAUSAL_CONV1D_RELEASE_TAG,
|
||||
release_base_url = (
|
||||
"https://github.com/Dao-AILab/causal-conv1d/releases/download"
|
||||
),
|
||||
)
|
||||
|
|
@ -883,7 +881,7 @@ def _install_fast_path_hooks(event_queue: Any) -> None:
|
|||
wrapped = _make_wrapper(original, install_fn, gate_name)
|
||||
setattr(_iu, gate_name, wrapped)
|
||||
rebound = _rebind_in_already_imported_modules(
|
||||
attr_name=gate_name, old_obj=original, new_obj=wrapped
|
||||
attr_name = gate_name, old_obj = original, new_obj = wrapped
|
||||
)
|
||||
rebound_total += rebound
|
||||
logger.info(
|
||||
|
|
|
|||
|
|
@ -630,24 +630,26 @@ def _patch_iu_gates(monkeypatch, fla_gate, conv_gate):
|
|||
|
||||
|
||||
def test_hook_installs_when_gate_returns_false(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=False)
|
||||
fla_gate = _make_fake_gate(initial_return = False)
|
||||
conv_gate = _make_fake_gate(initial_return = False)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock(side_effect=lambda eq: setattr(fla_gate, "next_return", True))
|
||||
tile_install = mock.Mock(side_effect=lambda eq: None)
|
||||
conv_install = mock.Mock(side_effect=lambda **kw: setattr(conv_gate, "next_return", True))
|
||||
fla_install = mock.Mock(
|
||||
side_effect = lambda eq: setattr(fla_gate, "next_return", True)
|
||||
)
|
||||
tile_install = mock.Mock(side_effect = lambda eq: None)
|
||||
conv_install = mock.Mock(
|
||||
side_effect = lambda **kw: setattr(conv_gate, "next_return", True)
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fla_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", tile_install
|
||||
)
|
||||
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", conv_install)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
|
|
@ -660,8 +662,8 @@ def test_hook_installs_when_gate_returns_false(monkeypatch):
|
|||
|
||||
|
||||
def test_hook_skips_install_when_gate_already_true(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=True)
|
||||
conv_gate = _make_fake_gate(initial_return=True)
|
||||
fla_gate = _make_fake_gate(initial_return = True)
|
||||
conv_gate = _make_fake_gate(initial_return = True)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock()
|
||||
|
|
@ -670,13 +672,11 @@ def test_hook_skips_install_when_gate_already_true(monkeypatch):
|
|||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fla_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", tile_install
|
||||
)
|
||||
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", conv_install)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
|
|
@ -688,23 +688,25 @@ def test_hook_skips_install_when_gate_already_true(monkeypatch):
|
|||
|
||||
|
||||
def test_hook_idempotent_on_repeat_call(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=False)
|
||||
fla_gate = _make_fake_gate(initial_return = False)
|
||||
conv_gate = _make_fake_gate(initial_return = False)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock(side_effect=lambda eq: setattr(fla_gate, "next_return", True))
|
||||
fla_install = mock.Mock(
|
||||
side_effect = lambda eq: setattr(fla_gate, "next_return", True)
|
||||
)
|
||||
tile_install = mock.Mock()
|
||||
conv_install = mock.Mock(side_effect=lambda **kw: setattr(conv_gate, "next_return", True))
|
||||
conv_install = mock.Mock(
|
||||
side_effect = lambda **kw: setattr(conv_gate, "next_return", True)
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fla_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", tile_install
|
||||
)
|
||||
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", conv_install)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
|
|
@ -718,8 +720,8 @@ def test_hook_idempotent_on_repeat_call(monkeypatch):
|
|||
|
||||
|
||||
def test_hook_handles_install_failure_gracefully(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=True) # bypass to focus on FLA
|
||||
fla_gate = _make_fake_gate(initial_return = False)
|
||||
conv_gate = _make_fake_gate(initial_return = True) # bypass to focus on FLA
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
def raising_install(eq):
|
||||
|
|
@ -732,9 +734,9 @@ def test_hook_handles_install_failure_gracefully(monkeypatch):
|
|||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: None
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: None)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
|
|
@ -743,8 +745,8 @@ def test_hook_handles_install_failure_gracefully(monkeypatch):
|
|||
|
||||
|
||||
def test_hook_can_be_disabled_via_env(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=False)
|
||||
fla_gate = _make_fake_gate(initial_return = False)
|
||||
conv_gate = _make_fake_gate(initial_return = False)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock()
|
||||
|
|
@ -753,7 +755,7 @@ def test_hook_can_be_disabled_via_env(monkeypatch):
|
|||
)
|
||||
monkeypatch.setenv(worker._FAST_PATH_HOOKS_SKIP_ENV, "1")
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
|
|
@ -764,8 +766,8 @@ def test_hook_can_be_disabled_via_env(monkeypatch):
|
|||
|
||||
|
||||
def test_hook_clears_lru_cache_before_first_check(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=True)
|
||||
conv_gate = _make_fake_gate(initial_return=True)
|
||||
fla_gate = _make_fake_gate(initial_return = True)
|
||||
conv_gate = _make_fake_gate(initial_return = True)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -775,9 +777,9 @@ def test_hook_clears_lru_cache_before_first_check(monkeypatch):
|
|||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: None
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: None)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
_iu.is_flash_linear_attention_available()
|
||||
|
|
@ -791,8 +793,8 @@ def test_hook_rewrites_previously_imported_module_bindings(monkeypatch):
|
|||
transformers.utils.import_utils alone does NOT reach those local
|
||||
bindings. The hook installer sweeps sys.modules and rebinds them.
|
||||
"""
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=True)
|
||||
fla_gate = _make_fake_gate(initial_return = False)
|
||||
conv_gate = _make_fake_gate(initial_return = True)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
# Create a fake modeling module that did `from ... import is_flash_linear_attention_available`.
|
||||
|
|
@ -811,9 +813,9 @@ def test_hook_rewrites_previously_imported_module_bindings(monkeypatch):
|
|||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: None
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: None)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
|
||||
# The fake module's local binding has been rewritten to the wrapper.
|
||||
assert fake_mod.is_flash_linear_attention_available is not fla_gate
|
||||
|
|
@ -834,10 +836,10 @@ def test_hook_skips_when_import_utils_unavailable(monkeypatch):
|
|||
return real_import(name, *a, **kw)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", fake_import)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
|
||||
# Should not raise.
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
worker._install_fast_path_hooks(event_queue = _FakeQueue())
|
||||
|
||||
|
||||
def test_substring_fallback_unchanged_when_hook_skipped(monkeypatch):
|
||||
|
|
@ -851,9 +853,13 @@ def test_substring_fallback_unchanged_when_hook_skipped(monkeypatch):
|
|||
monkeypatch.setenv(worker._FAST_PATH_HOOKS_SKIP_ENV, "1")
|
||||
|
||||
# Qwen3.5 model triggers install.
|
||||
worker._ensure_flash_linear_attention(event_queue=[], model_name="unsloth/Qwen3.5-2B")
|
||||
worker._ensure_flash_linear_attention(
|
||||
event_queue = [], model_name = "unsloth/Qwen3.5-2B"
|
||||
)
|
||||
assert install_mock.call_count == 1
|
||||
|
||||
# Llama doesn't.
|
||||
worker._ensure_flash_linear_attention(event_queue=[], model_name="meta-llama/Llama-3.1-8B")
|
||||
worker._ensure_flash_linear_attention(
|
||||
event_queue = [], model_name = "meta-llama/Llama-3.1-8B"
|
||||
)
|
||||
assert install_mock.call_count == 1
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue