diff --git a/studio/backend/utils/transformers_version.py b/studio/backend/utils/transformers_version.py index f1dd17e3bc..9bc9e1c94b 100644 --- a/studio/backend/utils/transformers_version.py +++ b/studio/backend/utils/transformers_version.py @@ -8,9 +8,13 @@ default 4.57.x that ships with Unsloth. When loading a LoRA adapter with a custom name, we resolve the base model from ``adapter_config.json`` and check *that* against the model list. -This module detects the model being loaded and ensures the correct transformers -version is installed before proceeding. The pip install is done *without* -``--force-reinstall`` to avoid rebuilding unrelated dependencies. +Strategy (sys.path overlay): + • The default transformers (4.57.x) always lives in site-packages. + • When 5.x is needed, we ``pip install --target --no-deps`` to a + separate directory and **prepend** it to ``sys.path``. + • To revert, we simply **remove** that directory from ``sys.path``. + • After either change we purge cached modules so the next import picks + up the correct version. """ import importlib @@ -18,12 +22,24 @@ import importlib.metadata import json import logging import os +import shutil import subprocess import sys +import tempfile from pathlib import Path logger = logging.getLogger(__name__) +# Ensure our logger is visible even if root logger isn't configured for INFO. +if not logger.handlers: + _handler = logging.StreamHandler() + _handler.setLevel(logging.INFO) + _handler.setFormatter( + logging.Formatter("[%(name)s|%(levelname)s]%(message)s") + ) + logger.addHandler(_handler) + logger.setLevel(logging.INFO) + # --------------------------------------------------------------------------- # Detection # --------------------------------------------------------------------------- @@ -41,14 +57,16 @@ TRANSFORMERS_5_MODEL_SUBSTRINGS: tuple[str, ...] = ( TRANSFORMERS_5_VERSION = "5.1.0" TRANSFORMERS_DEFAULT_VERSION = "4.57.1" +# Persistent directory for the transformers 5.x overlay +_OVERLAY_DIR = os.path.join(tempfile.gettempdir(), "transformers_5_overlay") + def _resolve_base_model(model_name: str) -> str: """If *model_name* points to a LoRA adapter, return its base model. - Checks for ``adapter_config.json`` locally first (covers the common case of - local output directories with custom names). For HF repo IDs that look - like they might be LoRA adapters, falls back to - ``get_base_model_from_lora``. + Checks for ``adapter_config.json`` locally first. Only calls the heavier + ``get_base_model_from_lora`` for paths that are actual local directories + (avoids noisy warnings for plain HF model IDs). Returns the original *model_name* unchanged if it is not a LoRA adapter. """ @@ -62,15 +80,14 @@ def _resolve_base_model(model_name: str) -> str: base = cfg.get("base_model_name_or_path") if base: logger.info( - "Resolved LoRA adapter '%s' → base model '%s'", model_name, base, + "Resolved LoRA adapter '%s' → base model '%s'", + model_name, base, ) return base except Exception as exc: logger.debug("Could not read %s: %s", adapter_cfg_path, exc) - # --- Only try the heavier fallback for paths that look like local dirs --- - # (Avoids triggering noisy warnings for plain HF model IDs like - # "unsloth/GLM-4.7-Flash" which are obviously not LoRA adapters.) + # --- Only try the heavier fallback for local directories ---------------- if local_path.is_dir(): try: from utils.models import get_base_model_from_lora @@ -84,7 +101,8 @@ def _resolve_base_model(model_name: str) -> str: return base except Exception as exc: logger.debug( - "get_base_model_from_lora failed for '%s': %s", model_name, exc, + "get_base_model_from_lora failed for '%s': %s", + model_name, exc, ) return model_name @@ -102,95 +120,111 @@ def needs_transformers_5(model_name: str) -> bool: # --------------------------------------------------------------------------- def _get_in_memory_version() -> str | None: - """Return the transformers version currently loaded in this process, - or None if transformers hasn't been imported yet.""" + """Return the transformers version currently loaded in this process.""" tf = sys.modules.get("transformers") if tf is not None: return getattr(tf, "__version__", None) return None -def _get_on_disk_version() -> str | None: - """Return the transformers version installed on disk (pip metadata).""" - try: - return importlib.metadata.version("transformers") - except importlib.metadata.PackageNotFoundError: - return None +# All top-level prefixes that hold references to transformers internals. +_PURGE_PREFIXES = ( + "transformers", + "unsloth", + "unsloth_zoo", + "peft", + "trl", + "accelerate", + "auto_gptq", + "bitsandbytes", +) -def _pip_install(package_spec: str) -> None: - """Run ``pip install `` using the running interpreter's pip. +def _purge_modules() -> int: + """Remove all cached modules for transformers and its dependents. - Mirrors the approach used in ``unsloth_zoo.llama_cpp.check_pip`` — - ``sys.executable -m pip`` is always the safest choice. - """ - cmd = [sys.executable, "-m", "pip", "install", package_spec] - logger.info("Running: %s", " ".join(cmd)) - result = subprocess.run( - cmd, - stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, - text=True, - ) - if result.returncode != 0: - logger.error("pip install failed:\n%s", result.stdout) - raise RuntimeError( - f"Failed to install {package_spec}. " - f"pip output:\n{result.stdout}" - ) - logger.info("pip install succeeded for %s", package_spec) - - -def _reload_transformers() -> None: - """Purge transformers AND all packages that cache transformers internals - from ``sys.modules``, then force a fresh re-import. - - Simply clearing ``transformers.*`` is not enough because libraries like - ``unsloth``, ``peft``, ``trl``, and ``accelerate`` bind transformers - classes/functions into their own module-level state. We must evict them - all so the next ``import`` picks up the freshly pip-installed version. + Returns the number of modules purged. """ importlib.invalidate_caches() - - # All top-level prefixes that hold references to transformers internals. - _PREFIXES = ( - "transformers", - "unsloth", - "unsloth_zoo", - "peft", - "trl", - "accelerate", - "auto_gptq", - "bitsandbytes", - ) - to_remove = [ k for k in list(sys.modules.keys()) - if any(k == p or k.startswith(p + ".") for p in _PREFIXES) + if any(k == p or k.startswith(p + ".") for p in _PURGE_PREFIXES) ] for key in to_remove: del sys.modules[key] + return len(to_remove) - logger.info("Purged %d cached modules (%s)", len(to_remove), - ", ".join(_PREFIXES)) - # Force a fresh import so the new version is loaded into the process. - import transformers # noqa: F811 - logger.info("Re-imported transformers — version is now %s", +def _install_overlay() -> None: + """Install transformers 5.x into the overlay directory and prepend + it to ``sys.path`` so it shadows the default site-packages version.""" + # Install if the overlay doesn't already exist + if not os.path.isdir(_OVERLAY_DIR) or not os.listdir(_OVERLAY_DIR): + os.makedirs(_OVERLAY_DIR, exist_ok=True) + cmd = [ + sys.executable, "-m", "pip", "install", + "--target", _OVERLAY_DIR, + "--no-deps", + f"transformers=={TRANSFORMERS_5_VERSION}", + ] + logger.info("Installing transformers %s to overlay: %s", + TRANSFORMERS_5_VERSION, " ".join(cmd)) + result = subprocess.run( + cmd, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + ) + if result.returncode != 0: + logger.error("pip install failed:\n%s", result.stdout) + raise RuntimeError( + f"Failed to install transformers=={TRANSFORMERS_5_VERSION} " + f"to {_OVERLAY_DIR}.\npip output:\n{result.stdout}" + ) + logger.info("Overlay install succeeded") + else: + logger.info("Overlay directory already exists at %s", _OVERLAY_DIR) + + # Prepend to sys.path (if not already there) + if _OVERLAY_DIR not in sys.path: + sys.path.insert(0, _OVERLAY_DIR) + logger.info("Prepended %s to sys.path", _OVERLAY_DIR) + + # Purge old modules and force fresh import + count = _purge_modules() + logger.info("Purged %d cached modules", count) + + import transformers + logger.info("Loaded transformers %s from overlay", transformers.__version__) + + +def _remove_overlay() -> None: + """Remove the overlay directory from ``sys.path`` so the default + site-packages version (4.57.x) takes effect again.""" + # Remove from sys.path + changed = False + while _OVERLAY_DIR in sys.path: + sys.path.remove(_OVERLAY_DIR) + changed = True + + if changed: + logger.info("Removed %s from sys.path", _OVERLAY_DIR) + + # Purge old modules and force fresh import + count = _purge_modules() + logger.info("Purged %d cached modules", count) + + import transformers + logger.info("Reverted to transformers %s from site-packages", transformers.__version__) def ensure_transformers_version(model_name: str) -> None: - """Ensure the correct ``transformers`` version is installed for *model_name*. + """Ensure the correct ``transformers`` version is active for *model_name*. - Checks BOTH the on-disk version (pip metadata) and the in-memory version - (``transformers.__version__``) because they can diverge after a previous - pip install in the same process. - - * If the model needs 5.x and the loaded version is already 5.x → no-op. - * If the model needs 5.x but 4.x is loaded → pip install 5.1.0 + reload. - * If the model does NOT need 5.x but 5.x is loaded → downgrade + reload. - * Otherwise → no-op. + Uses sys.path overlay: + • Need 5.x → install to separate dir, prepend sys.path, purge modules. + • Need 4.x → remove overlay from sys.path, purge modules. For LoRA adapters with custom names, the base model is resolved from ``adapter_config.json`` before checking. @@ -204,51 +238,33 @@ def ensure_transformers_version(model_name: str) -> None: target_version = TRANSFORMERS_5_VERSION if want_5 else TRANSFORMERS_DEFAULT_VERSION target_major = int(target_version.split(".")[0]) - # --- Check what's actually loaded in memory first ----------------------- + # Check what's actually loaded in memory in_memory = _get_in_memory_version() - on_disk = _get_on_disk_version() + overlay_active = _OVERLAY_DIR in sys.path logger.info( "Version check for '%s' (resolved: '%s'): need=%s, " - "in_memory=%s, on_disk=%s", - model_name, resolved, target_version, in_memory, on_disk, + "in_memory=%s, overlay_active=%s", + model_name, resolved, target_version, in_memory, overlay_active, ) + # --- Already correct? --------------------------------------------------- if in_memory is not None: in_memory_major = int(in_memory.split(".")[0]) if in_memory_major == target_major: logger.info( - "transformers %s in memory — already correct for '%s'", + "transformers %s already loaded — correct for '%s'", in_memory, model_name, ) return - # Wrong major in memory — need to switch - logger.info( - "transformers %s loaded in memory but need %s — switching…", - in_memory, target_version, - ) - elif on_disk is not None: - on_disk_major = int(on_disk.split(".")[0]) - if on_disk_major == target_major: - logger.info( - "transformers %s on disk (not yet imported) — correct for '%s'", - on_disk, model_name, - ) - return - logger.info( - "transformers %s on disk but need %s — switching…", - on_disk, target_version, - ) + + # --- Switch version ----------------------------------------------------- + if want_5: + logger.info("Activating transformers %s overlay…", TRANSFORMERS_5_VERSION) + _install_overlay() else: - logger.warning("transformers is not installed — skipping version check") - return - - # --- pip install the target version if needed --------------------------- - if on_disk is None or int(on_disk.split(".")[0]) != target_major: - _pip_install(f"transformers=={target_version}") - - # --- Purge and reload --------------------------------------------------- - _reload_transformers() + logger.info("Reverting to default transformers %s…", TRANSFORMERS_DEFAULT_VERSION) + _remove_overlay() final = _get_in_memory_version() - logger.info("Transformers version is now %s (in memory)", final) + logger.info("✓ transformers version is now %s", final)