Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
Next Next commit
Revert "feat: add CUDA device compatibility validation and correspond…
…ing tests"

This reverts commit 6d3e514.
  • Loading branch information
umisetokikaze committed Mar 11, 2026
commit d160880b709f7146e7c55cf9e52e77c434c7ce5e
43 changes: 0 additions & 43 deletions library/device_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,49 +87,6 @@ def get_preferred_device() -> torch.device:
return device



def _normalize_cuda_arch(arch) -> Optional[str]:
if isinstance(arch, str):
return arch if arch.startswith("sm_") else None
if isinstance(arch, (tuple, list)) and len(arch) >= 2:
return f"sm_{int(arch[0])}{int(arch[1])}"
return None


def validate_cuda_device_compatibility(device: Optional[Union[str, torch.device]] = None):
if not HAS_CUDA:
return

if device is None:
device = torch.device("cuda")
elif isinstance(device, str):
device = torch.device(device)

if device.type != "cuda":
return

get_arch_list = getattr(torch.cuda, "get_arch_list", None)
if get_arch_list is None:
return

try:
supported_arches = sorted(
{arch_name for arch_name in (_normalize_cuda_arch(arch) for arch in get_arch_list()) if arch_name is not None}
)
device_arch = _normalize_cuda_arch(torch.cuda.get_device_capability(device))
device_name = torch.cuda.get_device_name(device)
except Exception:
return

if supported_arches and device_arch is not None and device_arch not in supported_arches:
cuda_version = getattr(torch.version, "cuda", None)
cuda_suffix = f" with CUDA {cuda_version}" if cuda_version else ""
supported = ", ".join(supported_arches)
raise RuntimeError(
f"CUDA device '{device_name}' reports {device_arch}, but this PyTorch build{cuda_suffix} only supports {supported}. "
+ "Install a PyTorch build that includes kernels for this GPU from https://pytorch.org/get-started/locally/ or build PyTorch from source."
)

def init_ipex():
"""
Apply IPEX to CUDA hijacks using `library.ipex.ipex_init`.
Expand Down
3 changes: 1 addition & 2 deletions library/train_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@
from packaging.version import Version

import torch
from library.device_utils import init_ipex, clean_memory_on_device, validate_cuda_device_compatibility
from library.device_utils import init_ipex, clean_memory_on_device
from library.strategy_base import LatentsCachingStrategy, TokenizeStrategy, TextEncoderOutputsCachingStrategy, TextEncodingStrategy

init_ipex()
Expand Down Expand Up @@ -5500,7 +5500,6 @@ def prepare_accelerator(args: argparse.Namespace):
dynamo_backend=dynamo_backend,
deepspeed_plugin=deepspeed_plugin,
)
validate_cuda_device_compatibility(accelerator.device)
print("accelerator device:", accelerator.device)
return accelerator

Expand Down
24 changes: 0 additions & 24 deletions tests/test_device_utils.py

This file was deleted.