From d2d758d0d8088cc1af46f820a481223efd9d4087 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sat, 16 May 2026 12:55:51 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/training/worker.py | 22 ++--- .../tests/test_training_worker_flash_attn.py | 96 ++++++++++--------- 2 files changed, 61 insertions(+), 57 deletions(-) diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index d717a3a5e9..b65a21568c 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -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( diff --git a/studio/backend/tests/test_training_worker_flash_attn.py b/studio/backend/tests/test_training_worker_flash_attn.py index d8802864c2..ddd66a4fdc 100644 --- a/studio/backend/tests/test_training_worker_flash_attn.py +++ b/studio/backend/tests/test_training_worker_flash_attn.py @@ -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