Compare commits

...
Sign in to create a new pull request.

1 commit

Author SHA1 Message Date
Kaitao Yang
05354af3a6 introduce device_context to simplify code. 2026-03-12 10:10:01 +00:00
3 changed files with 71 additions and 73 deletions

View file

@ -20,6 +20,10 @@ __all__ = [
"DEVICE_COUNT",
"ALLOW_PREQUANTIZED_MODELS",
"ALLOW_BITSANDBYTES",
"DeviceContext",
"device_context",
"clean_gpu_cache",
"get_current_device",
]
import torch
@ -130,3 +134,57 @@ if DEVICE_TYPE == "hip":
Params4bit
):
ALLOW_PREQUANTIZED_MODELS = False
class DeviceContext:
"""Encapsulates device-specific operations for XPU/HIP/CUDA."""
def __init__(self, device_type: str = DEVICE_TYPE) -> None:
DEVICE_MODULE_MAP = {"xpu": torch.xpu, "cuda": torch.cuda, "hip": torch.cuda}
if device_type not in DEVICE_MODULE_MAP:
raise ValueError(f"Unsloth: Unsupported device type: {device_type}")
self.device_type = device_type
# Cache the torch module for this device
self.torch_module = DEVICE_MODULE_MAP[device_type]
def get_stats(self) -> tuple[str, str, float]:
"""Return (name, stats_snippet, max_memory_gb)."""
gpu_stats = self.torch_module.get_device_properties(0)
max_mem = round(gpu_stats.total_memory / 1024 / 1024 / 1024, 3)
# Device name
name = gpu_stats.name + ". " if gpu_stats.name else self._get_default_name()
# Toolkit snippet
snippet = self._get_toolkit_snippet(gpu_stats)
return name, snippet, max_mem
def _get_default_name(self) -> str:
"""Get default device name when props.name is empty."""
names = {"xpu": "Intel XPU", "cuda": "NVIDIA GPU", "hip": "AMD GPU"}
return names[self.device_type] + " Device. "
def _get_toolkit_snippet(self, props) -> str:
"""Get toolkit version snippet."""
if self.device_type == "cuda":
return f"CUDA: {props.major}.{props.minor}. CUDA Toolkit: {torch.version.cuda}."
elif self.device_type == "hip":
return f"ROCm Toolkit: {torch.version.hip}."
else: # xpu
return f"Intel Toolkit: {torch.version.xpu}."
# Singleton instance
device_context = DeviceContext()
# Module-level functions for backward compatibility
def clean_gpu_cache() -> None:
"""Clear GPU cache for current device type."""
device_context.torch_module.empty_cache()
def get_current_device() -> int:
"""Get current device index."""
return device_context.torch_module.current_device()

View file

@ -56,6 +56,9 @@ from ..device_type import (
DEVICE_TYPE_TORCH,
DEVICE_COUNT,
ALLOW_PREQUANTIZED_MODELS,
device_context,
clean_gpu_cache,
get_current_device,
)
transformers_version = Version(transformers_version)
@ -123,13 +126,6 @@ BlockDiagonalCausalMask = (
xformers.attn_bias.BlockDiagonalCausalMask if HAS_XFORMERS else None
)
if DEVICE_TYPE == "xpu":
clean_gpu_cache = torch.xpu.empty_cache
get_current_device = torch.xpu.current_device
else:
clean_gpu_cache = torch.cuda.empty_cache
get_current_device = torch.cuda.current_device
def original_apply_qkv(self, X):
Q = self.q_proj(X)
@ -2245,41 +2241,12 @@ class FastLlamaModel:
model_patcher = FastLlamaModel
SUPPORTS_BFLOAT16 = is_bfloat16_supported()
if DEVICE_TYPE == "cuda":
gpu_stats = torch.cuda.get_device_properties(0)
gpu_stats_name = (
gpu_stats.name + ". " if gpu_stats.name != "" else "NVIDIA GPU Device. "
)
gpu_version = torch.version.cuda
gpu_stats_snippet = f"CUDA: {gpu_stats.major}.{gpu_stats.minor}. CUDA Toolkit: {gpu_version}."
try:
vllm_version = f" vLLM: {importlib_version('vllm')}."
except:
vllm_version = ""
elif DEVICE_TYPE == "hip":
gpu_stats = torch.cuda.get_device_properties(0)
gpu_stats_name = resolve_hip_gpu_stats_name(gpu_stats)
gpu_version = torch.version.hip
gpu_stats_snippet = f"ROCm Toolkit: {gpu_version}."
try:
vllm_version = f" vLLM: {importlib_version('vllm')}."
except:
vllm_version = ""
elif DEVICE_TYPE == "xpu":
gpu_stats = torch.xpu.get_device_properties(0)
gpu_stats_name = (
gpu_stats.name + ". " if gpu_stats.name != "" else "Intel XPU Device. "
)
gpu_version = torch.version.xpu
gpu_stats_snippet = f"Intel Toolkit: {gpu_version}."
try:
vllm_version = f" vLLM: {importlib_version('vllm')}."
except:
vllm_version = ""
else:
raise ValueError(f"Unsloth: Unsupported device type: {DEVICE_TYPE}")
gpu_stats_name, gpu_stats_snippet, max_memory = device_context.get_stats()
max_memory = round(gpu_stats.total_memory / 1024 / 1024 / 1024, 3)
try:
vllm_version = f" vLLM: {importlib_version('vllm')}."
except:
vllm_version = ""
statistics = (
f"==((====))== Unsloth {__version__}: Fast {model_patcher.__name__[4:-5]} patching. Transformers: {transformers_version}.{vllm_version}\n"

View file

@ -79,6 +79,7 @@ except:
# Old HF Hub versions <= 0.0.25
from huggingface_hub.utils._token import get_token
from ..device_type import (
device_context,
is_hip,
get_device_type,
DEVICE_TYPE,
@ -479,39 +480,11 @@ class FastBaseModel:
token = hf_login(token)
SUPPORTS_BFLOAT16 = is_bfloat16_supported()
if DEVICE_TYPE == "cuda":
gpu_stats = torch.cuda.get_device_properties(0)
gpu_stats_name = (
gpu_stats.name + ". " if gpu_stats.name != "" else "NVIDIA GPU Device. "
)
gpu_version = torch.version.cuda
gpu_stats_snippet = f"CUDA: {gpu_stats.major}.{gpu_stats.minor}. CUDA Toolkit: {gpu_version}."
try:
vllm_version = f" vLLM: {importlib_version('vllm')}."
except:
vllm_version = ""
elif DEVICE_TYPE == "hip":
gpu_stats = torch.cuda.get_device_properties(0)
gpu_stats_name = resolve_hip_gpu_stats_name(gpu_stats)
gpu_version = torch.version.hip
gpu_stats_snippet = f"ROCm Toolkit: {gpu_version}."
try:
vllm_version = f" vLLM: {importlib_version('vllm')}."
except:
vllm_version = ""
elif DEVICE_TYPE == "xpu":
gpu_stats = torch.xpu.get_device_properties(0)
gpu_stats_name = (
gpu_stats.name + ". " if gpu_stats.name != "" else "Intel XPU Device. "
)
gpu_version = torch.version.xpu
gpu_stats_snippet = f"Intel Toolkit: {gpu_version}."
# [TODO] After adding vLLM support for XPU, change this
gpu_stats_name, gpu_stats_snippet, max_memory = device_context.get_stats()
try:
vllm_version = f" vLLM: {importlib_version('vllm')}."
except:
vllm_version = ""
else:
raise ValueError(f"Unsloth: Unsupported device type: {DEVICE_TYPE}")
max_memory = round(gpu_stats.total_memory / 1024 / 1024 / 1024, 3)
arch_name = model_type_arch.title()
arch_name = arch_name.replace("_Vl_", "_VL_").replace("_Moe", "_MoE")