unsloth/studio/backend/utils/utils.py
2026-02-11 16:04:35 +00:00

106 lines
3.3 KiB
Python

"""
Shared backend utilities
"""
import os
import logging
from contextlib import contextmanager
from pathlib import Path
import shutil
import tempfile
logger = logging.getLogger(__name__)
@contextmanager
def without_hf_auth():
"""
Context manager to temporarily disable HuggingFace authentication.
Usage:
with without_hf_auth():
# Code that should run without cached tokens
model_info(model_name, token=None)
"""
# Save environment variables
saved_env = {}
env_vars = ['HF_TOKEN', 'HUGGINGFACE_HUB_TOKEN', 'HF_HOME']
for var in env_vars:
if var in os.environ:
saved_env[var] = os.environ[var]
del os.environ[var]
# Save disable flag
saved_disable = os.environ.get('HF_HUB_DISABLE_IMPLICIT_TOKEN')
os.environ['HF_HUB_DISABLE_IMPLICIT_TOKEN'] = '1'
# Move token files temporarily
token_files = []
token_locations = [
Path.home() / '.cache' / 'huggingface' / 'token',
Path.home() / '.huggingface' / 'token'
]
for token_loc in token_locations:
if token_loc.exists():
temp = tempfile.NamedTemporaryFile(delete=False)
temp.close()
shutil.move(str(token_loc), temp.name)
token_files.append((token_loc, temp.name))
try:
yield
finally:
# Restore tokens
for original, temp in token_files:
try:
original.parent.mkdir(parents=True, exist_ok=True)
shutil.move(temp, str(original))
except Exception as e:
logger.error(f"Failed to restore token {original}: {e}")
# Restore environment
for var, value in saved_env.items():
os.environ[var] = value
if saved_disable is not None:
os.environ['HF_HUB_DISABLE_IMPLICIT_TOKEN'] = saved_disable
else:
os.environ.pop('HF_HUB_DISABLE_IMPLICIT_TOKEN', None)
pass
def format_error_message(error: Exception, model_name: str) -> str:
"""
Format user-friendly error messages for common issues.
Args:
error: The exception that occurred
model_name: Name of the model being loaded
Returns:
User-friendly error string
"""
error_str = str(error).lower()
model_short = model_name.split('/')[-1] if '/' in model_name else model_name
if "repository not found" in error_str or "404" in error_str:
return f"Model '{model_short}' not found. Check the model name."
if "401" in error_str or "unauthorized" in error_str:
return f"Authentication failed for '{model_short}'. Please provide a valid HF token."
if "gated" in error_str or "access to model" in error_str:
return f"Model '{model_short}' requires authentication. Please provide a valid HF token."
if "invalid user token" in error_str:
return "Invalid HF token. Please check your token and try again."
if "memory" in error_str or "cuda" in error_str or "mlx" in error_str or "out of memory" in error_str:
from utils.hardware import get_device
device = get_device()
device_label = {"cuda": "GPU", "mlx": "Apple Silicon GPU", "cpu": "system"}.get(device.value, "GPU")
return f"Not enough {device_label} memory to load '{model_short}'. Try a smaller model or free memory."
# Generic fallback
return str(error)
pass