extract shared activate_transformers_for_subprocess into transformers_version.py
This commit is contained in:
parent
532d360094
commit
41fce49d6f
4 changed files with 52 additions and 89 deletions
|
|
@ -181,8 +181,11 @@ def _setup_log_capture(resp_queue: Any) -> None:
|
|||
|
||||
def _activate_transformers_version(model_name: str) -> None:
|
||||
"""Activate the correct transformers version BEFORE any ML imports."""
|
||||
# Ensure backend is on path for utils imports
|
||||
backend_path = str(Path(__file__).resolve().parent.parent.parent)
|
||||
if backend_path not in sys.path:
|
||||
sys.path.insert(0, backend_path)
|
||||
|
||||
from utils.transformers_version import activate_transformers_for_subprocess
|
||||
|
||||
activate_transformers_for_subprocess(model_name)
|
||||
|
|
@ -203,19 +206,6 @@ def _handle_load(backend, cmd: dict, resp_queue: Any) -> None:
|
|||
load_in_4bit = cmd.get("load_in_4bit", True)
|
||||
trust_remote_code = cmd.get("trust_remote_code", False)
|
||||
|
||||
# Auto-enable trust_remote_code for NemotronH/Nano models.
|
||||
if not trust_remote_code:
|
||||
_NEMOTRON_TRUST_SUBSTRINGS = ("nemotron_h", "nemotron-h", "nemotron-3-nano")
|
||||
_cp_lower = checkpoint_path.lower()
|
||||
if any(sub in _cp_lower for sub in _NEMOTRON_TRUST_SUBSTRINGS) and (
|
||||
_cp_lower.startswith("unsloth/") or _cp_lower.startswith("nvidia/")
|
||||
):
|
||||
trust_remote_code = True
|
||||
logger.info(
|
||||
"Auto-enabled trust_remote_code for Nemotron model: %s",
|
||||
checkpoint_path,
|
||||
)
|
||||
|
||||
try:
|
||||
_send_response(
|
||||
resp_queue,
|
||||
|
|
|
|||
|
|
@ -34,50 +34,15 @@ from utils.hardware import apply_gpu_ids
|
|||
|
||||
|
||||
def _activate_transformers_version(model_name: str) -> None:
|
||||
"""Activate the correct transformers version BEFORE any ML imports.
|
||||
|
||||
Uses get_transformers_tier() to decide between .venv_t5_550/ (5.5.0),
|
||||
.venv_t5_530/ (5.3.0), or the default 4.57.x.
|
||||
"""
|
||||
"""Activate the correct transformers version BEFORE any ML imports."""
|
||||
# Ensure backend is on path for utils imports
|
||||
backend_path = str(Path(__file__).resolve().parent.parent.parent)
|
||||
if backend_path not in sys.path:
|
||||
sys.path.insert(0, backend_path)
|
||||
|
||||
from utils.transformers_version import (
|
||||
get_transformers_tier,
|
||||
_resolve_base_model,
|
||||
_ensure_venv_t5_530_exists,
|
||||
_ensure_venv_t5_550_exists,
|
||||
_VENV_T5_530_DIR,
|
||||
_VENV_T5_550_DIR,
|
||||
)
|
||||
from utils.transformers_version import activate_transformers_for_subprocess
|
||||
|
||||
resolved = _resolve_base_model(model_name)
|
||||
tier = get_transformers_tier(resolved)
|
||||
|
||||
if tier == "550":
|
||||
if not _ensure_venv_t5_550_exists():
|
||||
raise RuntimeError(
|
||||
f"Cannot activate transformers 5.5.0: .venv_t5_550 missing at {_VENV_T5_550_DIR}"
|
||||
)
|
||||
if _VENV_T5_550_DIR not in sys.path:
|
||||
sys.path.insert(0, _VENV_T5_550_DIR)
|
||||
logger.info("Activated transformers 5.5.0 from %s", _VENV_T5_550_DIR)
|
||||
_pp = os.environ.get("PYTHONPATH", "")
|
||||
os.environ["PYTHONPATH"] = _VENV_T5_550_DIR + (os.pathsep + _pp if _pp else "")
|
||||
elif tier == "530":
|
||||
if not _ensure_venv_t5_530_exists():
|
||||
raise RuntimeError(
|
||||
f"Cannot activate transformers 5.3.0: .venv_t5_530 missing at {_VENV_T5_530_DIR}"
|
||||
)
|
||||
if _VENV_T5_530_DIR not in sys.path:
|
||||
sys.path.insert(0, _VENV_T5_530_DIR)
|
||||
logger.info("Activated transformers 5.3.0 from %s", _VENV_T5_530_DIR)
|
||||
_pp = os.environ.get("PYTHONPATH", "")
|
||||
os.environ["PYTHONPATH"] = _VENV_T5_530_DIR + (os.pathsep + _pp if _pp else "")
|
||||
else:
|
||||
logger.info("Using default transformers (4.57.x) for %s", model_name)
|
||||
activate_transformers_for_subprocess(model_name)
|
||||
|
||||
|
||||
def _decode_image(image_base64: str):
|
||||
|
|
|
|||
|
|
@ -1039,50 +1039,15 @@ def _ensure_flash_attn_for_long_context(event_queue: Any, max_seq_length: int) -
|
|||
|
||||
|
||||
def _activate_transformers_version(model_name: str) -> None:
|
||||
"""Activate the correct transformers version BEFORE any ML imports.
|
||||
|
||||
Uses get_transformers_tier() to decide between .venv_t5_550/ (5.5.0),
|
||||
.venv_t5_530/ (5.3.0), or the default 4.57.x.
|
||||
"""
|
||||
"""Activate the correct transformers version BEFORE any ML imports."""
|
||||
# Ensure backend is on path for utils imports
|
||||
backend_path = str(Path(__file__).resolve().parent.parent.parent)
|
||||
if backend_path not in sys.path:
|
||||
sys.path.insert(0, backend_path)
|
||||
|
||||
from utils.transformers_version import (
|
||||
get_transformers_tier,
|
||||
_resolve_base_model,
|
||||
_ensure_venv_t5_530_exists,
|
||||
_ensure_venv_t5_550_exists,
|
||||
_VENV_T5_530_DIR,
|
||||
_VENV_T5_550_DIR,
|
||||
)
|
||||
from utils.transformers_version import activate_transformers_for_subprocess
|
||||
|
||||
resolved = _resolve_base_model(model_name)
|
||||
tier = get_transformers_tier(resolved)
|
||||
|
||||
if tier == "550":
|
||||
if not _ensure_venv_t5_550_exists():
|
||||
raise RuntimeError(
|
||||
f"Cannot activate transformers 5.5.0: .venv_t5_550 missing at {_VENV_T5_550_DIR}"
|
||||
)
|
||||
if _VENV_T5_550_DIR not in sys.path:
|
||||
sys.path.insert(0, _VENV_T5_550_DIR)
|
||||
logger.info("Activated transformers 5.5.0 from %s", _VENV_T5_550_DIR)
|
||||
_pp = os.environ.get("PYTHONPATH", "")
|
||||
os.environ["PYTHONPATH"] = _VENV_T5_550_DIR + (os.pathsep + _pp if _pp else "")
|
||||
elif tier == "530":
|
||||
if not _ensure_venv_t5_530_exists():
|
||||
raise RuntimeError(
|
||||
f"Cannot activate transformers 5.3.0: .venv_t5_530 missing at {_VENV_T5_530_DIR}"
|
||||
)
|
||||
if _VENV_T5_530_DIR not in sys.path:
|
||||
sys.path.insert(0, _VENV_T5_530_DIR)
|
||||
logger.info("Activated transformers 5.3.0 from %s", _VENV_T5_530_DIR)
|
||||
_pp = os.environ.get("PYTHONPATH", "")
|
||||
os.environ["PYTHONPATH"] = _VENV_T5_530_DIR + (os.pathsep + _pp if _pp else "")
|
||||
else:
|
||||
logger.info("Using default transformers (4.57.x) for %s", model_name)
|
||||
activate_transformers_for_subprocess(model_name)
|
||||
|
||||
|
||||
def _adapt_for_mlx_vlm(items):
|
||||
|
|
|
|||
|
|
@ -112,6 +112,49 @@ _VENV_T5_550_DIR = str(_studio_root() / ".venv_t5_550")
|
|||
_VENV_T5_DIR = _VENV_T5_550_DIR
|
||||
|
||||
|
||||
def activate_transformers_for_subprocess(model_name: str) -> None:
|
||||
"""Activate the correct transformers version in a subprocess worker.
|
||||
|
||||
Call this BEFORE any ML imports. Resolves LoRA adapters to their base
|
||||
model, determines the required tier, and prepends the appropriate
|
||||
``.venv_t5_*`` directory to ``sys.path``. Also propagates the path
|
||||
via ``PYTHONPATH`` for child processes (e.g. GGUF converter).
|
||||
|
||||
Used by training, inference, and export workers.
|
||||
"""
|
||||
resolved = _resolve_base_model(model_name)
|
||||
tier = get_transformers_tier(resolved)
|
||||
|
||||
if tier == "550":
|
||||
if not _ensure_venv_t5_550_exists():
|
||||
raise RuntimeError(
|
||||
f"Cannot activate transformers 5.5.0: "
|
||||
f".venv_t5_550 missing at {_VENV_T5_550_DIR}"
|
||||
)
|
||||
if _VENV_T5_550_DIR not in sys.path:
|
||||
sys.path.insert(0, _VENV_T5_550_DIR)
|
||||
logger.info("Activated transformers 5.5.0 from %s", _VENV_T5_550_DIR)
|
||||
_pp = os.environ.get("PYTHONPATH", "")
|
||||
os.environ["PYTHONPATH"] = (
|
||||
_VENV_T5_550_DIR + (os.pathsep + _pp if _pp else "")
|
||||
)
|
||||
elif tier == "530":
|
||||
if not _ensure_venv_t5_530_exists():
|
||||
raise RuntimeError(
|
||||
f"Cannot activate transformers 5.3.0: "
|
||||
f".venv_t5_530 missing at {_VENV_T5_530_DIR}"
|
||||
)
|
||||
if _VENV_T5_530_DIR not in sys.path:
|
||||
sys.path.insert(0, _VENV_T5_530_DIR)
|
||||
logger.info("Activated transformers 5.3.0 from %s", _VENV_T5_530_DIR)
|
||||
_pp = os.environ.get("PYTHONPATH", "")
|
||||
os.environ["PYTHONPATH"] = (
|
||||
_VENV_T5_530_DIR + (os.pathsep + _pp if _pp else "")
|
||||
)
|
||||
else:
|
||||
logger.info("Using default transformers (4.57.x) for %s", model_name)
|
||||
|
||||
|
||||
def _resolve_base_model(model_name: str) -> str:
|
||||
"""If *model_name* points to a LoRA adapter, return its base model.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue