mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
get_hip_arch_name -> get_hip_agent, use amdhip64_7.dll served within rocm package
This commit is contained in:
+1
-1
@@ -946,7 +946,7 @@ def check_torch():
|
||||
|
||||
if not is_cuda_available and not is_ipex_available and allow_rocm:
|
||||
from modules import rocm
|
||||
is_rocm_available = allow_rocm and (args.use_rocm or args.use_zluda or (len(rocm.agents) != 0 if sys.platform == "win32" else rocm.is_installed)) # late eval to avoid unnecessary import
|
||||
is_rocm_available = allow_rocm and (args.use_rocm or args.use_zluda or rocm.is_installed) # late eval to avoid unnecessary import
|
||||
|
||||
if is_cuda_available and args.use_cuda: # prioritize cuda
|
||||
torch_command = install_cuda()
|
||||
|
||||
@@ -89,7 +89,7 @@ def set_ck_flash_attention(backend: str, device: torch.device):
|
||||
if backend == "rocm":
|
||||
if not installed('flash-attn'):
|
||||
log.info('Torch attention: type="Flash attention" building...')
|
||||
agent = rocm.Agent(getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000"))
|
||||
agent = rocm.Agent(device)
|
||||
install(rocm.get_flash_attention_command(agent), reinstall=True)
|
||||
else:
|
||||
install('flash-attn')
|
||||
|
||||
+5
-7
@@ -67,10 +67,8 @@ def has_triton(early:bool=False) -> bool:
|
||||
return test_triton(early=early)
|
||||
|
||||
|
||||
def get_hip_arch_name() -> str:
|
||||
arch_name = getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000")
|
||||
arch_name = arch_name.split(':')[0]
|
||||
return arch_name
|
||||
def get_hip_agent() -> rocm.Agent:
|
||||
return rocm.Agent(device)
|
||||
|
||||
|
||||
def get_backend(shared_cmd_opts):
|
||||
@@ -330,9 +328,9 @@ def test_fp16():
|
||||
elif backend == 'rocm':
|
||||
# gfx1102 (RX 7600, 7500, 7650 and 7700S) causes segfaults with fp16
|
||||
# agent can be overriden to gfx1100 to get gfx1102 working with ROCm so check the gpu name as well
|
||||
agent = get_hip_arch_name()
|
||||
agent = get_hip_agent()
|
||||
agent_name = getattr(torch.cuda.get_device_properties(device), "name", "AMD Radeon RX 0000")
|
||||
if agent == "gfx1102" or (agent == "gfx1100" and any(i in agent_name for i in ("7600", "7500", "7650", "7700S"))):
|
||||
if agent.gfx_version == 0x1102 or (agent.gfx_version == 0x1100 and any(i in agent_name for i in ("7600", "7500", "7650", "7700S"))):
|
||||
fp16_ok = False
|
||||
return fp16_ok
|
||||
try:
|
||||
@@ -361,7 +359,7 @@ def test_bf16():
|
||||
elif backend == 'rocm' or backend == 'zluda':
|
||||
agent = None
|
||||
if backend == 'rocm':
|
||||
agent = rocm.Agent(get_hip_arch_name())
|
||||
agent = get_hip_agent()
|
||||
else:
|
||||
from modules.zluda_installer import default_agent
|
||||
agent = default_agent
|
||||
|
||||
+128
-37
@@ -3,9 +3,11 @@ import sys
|
||||
import ctypes
|
||||
import shutil
|
||||
import subprocess
|
||||
from typing import Union, List
|
||||
from typing import Union, overload, TYPE_CHECKING
|
||||
from enum import Enum
|
||||
from functools import wraps
|
||||
if TYPE_CHECKING:
|
||||
import torch
|
||||
|
||||
|
||||
def resolve_link(path_: str) -> str:
|
||||
@@ -20,8 +22,8 @@ def dirname(path_: str, r: int = 1) -> str:
|
||||
return path_
|
||||
|
||||
|
||||
def spawn(command: Union[str, List[str]], cwd: os.PathLike = '.') -> str:
|
||||
process = subprocess.run(command, cwd=cwd, shell=True, check=False, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
|
||||
def spawn(command: Union[str, list[str]], cwd: os.PathLike = '.') -> str:
|
||||
process = subprocess.run(command, cwd=cwd, shell=True, check=False, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL)
|
||||
return process.stdout.decode(encoding="utf8", errors="ignore")
|
||||
|
||||
|
||||
@@ -83,13 +85,18 @@ class Agent:
|
||||
break
|
||||
return result
|
||||
|
||||
def __init__(self, name: str = None):
|
||||
if name is None:
|
||||
@overload
|
||||
def __init__(self, name: str): ...
|
||||
@overload
|
||||
def __init__(self, device: 'torch.types.Device'): ...
|
||||
|
||||
def __init__(self, arg):
|
||||
if isinstance(arg, str):
|
||||
name = arg
|
||||
else: # assume arg is device-like object
|
||||
import torch
|
||||
device = torch.cuda.current_device()
|
||||
name = getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000").split(':')[0]
|
||||
else:
|
||||
self.name = name.split(':')[0]
|
||||
name = getattr(torch.cuda.get_device_properties(arg), "gcnArchName", "gfx0000")
|
||||
self.name = name.split(':')[0]
|
||||
self.gfx_version = Agent.parse_gfx_version(self.name)
|
||||
if self.gfx_version > 0x1000:
|
||||
self.arch = MicroArchitecture.RDNA
|
||||
@@ -224,12 +231,6 @@ def get_flash_attention_command(agent: Agent) -> str:
|
||||
|
||||
def refresh():
|
||||
global environment, blaslt_tensile_libpath, is_installed, version # pylint: disable=global-statement
|
||||
if sys.platform == "win32":
|
||||
global agents # pylint: disable=global-statement
|
||||
try:
|
||||
agents = driver_get_agents()
|
||||
except Exception:
|
||||
agents = []
|
||||
environment = find()
|
||||
if environment is not None:
|
||||
if isinstance(environment, ROCmEnvironment):
|
||||
@@ -241,24 +242,20 @@ def refresh():
|
||||
|
||||
|
||||
if sys.platform == "win32":
|
||||
def get_agents() -> List[Agent]:
|
||||
return agents
|
||||
#if isinstance(environment, ROCmEnvironment):
|
||||
# out = spawn("amdgpu-arch", cwd=os.path.join(environment.path, 'bin'))
|
||||
#else:
|
||||
# # Assume that amdgpu-arch is in PATH (venv/Scripts/amdgpu-arch.exe)
|
||||
# out = spawn("amdgpu-arch")
|
||||
#out = out.strip()
|
||||
#if out == "":
|
||||
# return []
|
||||
#return [Agent(x.split(' ')[-1].strip()) for x in out.split("\n")]
|
||||
import tempfile
|
||||
|
||||
def driver_get_agents() -> List[Agent]:
|
||||
# unsafe and experimental feature
|
||||
from modules import windows_hip_ffi
|
||||
archs = windows_hip_ffi.get_archs()
|
||||
# filter out None (is there any better way?)
|
||||
return [Agent(x) for x in archs if x is not None]
|
||||
def get_agents() -> list[Agent]:
|
||||
name = None
|
||||
with tempfile.NamedTemporaryFile(mode='w', encoding='utf-8', delete=False) as f:
|
||||
name = f.name
|
||||
f.write(CODE_AMDGPU_ARCH)
|
||||
f.flush()
|
||||
out = spawn([sys.executable, name])
|
||||
os.unlink(name)
|
||||
out = out.strip()
|
||||
if out == "":
|
||||
return []
|
||||
return [Agent(x.split(' ')[-1].strip()) for x in out.split("\n")]
|
||||
|
||||
def postinstall():
|
||||
import torch
|
||||
@@ -274,6 +271,7 @@ if sys.platform == "win32":
|
||||
return
|
||||
|
||||
build_targets = torch.cuda.get_arch_list()
|
||||
agents = get_agents()
|
||||
for available in agents:
|
||||
if available.name in build_targets:
|
||||
return
|
||||
@@ -285,10 +283,10 @@ if sys.platform == "win32":
|
||||
try:
|
||||
import torch
|
||||
import numpy as np
|
||||
from modules.devices import get_hip_arch_name
|
||||
from modules.devices import get_hip_agent
|
||||
|
||||
gfx_version = Agent.parse_gfx_version(get_hip_arch_name())
|
||||
if (gfx_version & 0xFFF0) == 0x1200:
|
||||
agent = get_hip_agent()
|
||||
if (agent.gfx_version & 0xFFF0) == 0x1200:
|
||||
# disable MIOpen for gfx120x
|
||||
torch.backends.cudnn.enabled = False
|
||||
|
||||
@@ -319,9 +317,8 @@ if sys.platform == "win32":
|
||||
return True, None
|
||||
|
||||
is_wsl: bool = False
|
||||
agents: List[Agent] = [] # temp
|
||||
else: # sys.platform != "win32"
|
||||
def get_agents() -> List[Agent]:
|
||||
def get_agents() -> list[Agent]:
|
||||
try:
|
||||
_agents = spawn("rocm_agent_enumerator").split("\n")
|
||||
_agents = [x for x in _agents if x and x != 'gfx000']
|
||||
@@ -351,3 +348,97 @@ blaslt_tensile_libpath = ""
|
||||
is_installed = False
|
||||
version = None
|
||||
refresh()
|
||||
|
||||
# amdgpu-arch.exe written in Python
|
||||
CODE_AMDGPU_ARCH = """
|
||||
import sys
|
||||
if sys.platform == "win32":
|
||||
import os
|
||||
import ctypes
|
||||
import ctypes.wintypes
|
||||
import contextlib
|
||||
hipDeviceProp = ctypes.c_byte * 1472
|
||||
@contextlib.contextmanager
|
||||
def mute(fd):
|
||||
s = os.dup(fd)
|
||||
try:
|
||||
with open(os.devnull, 'w') as devnull:
|
||||
os.dup2(devnull.fileno(), fd)
|
||||
yield
|
||||
finally:
|
||||
os.dup2(s, fd)
|
||||
os.close(s)
|
||||
class HIP:
|
||||
def __init__(self):
|
||||
ctypes.windll.kernel32.LoadLibraryA.restype = ctypes.wintypes.HMODULE
|
||||
ctypes.windll.kernel32.LoadLibraryA.argtypes = [ctypes.c_char_p]
|
||||
self.handle = None
|
||||
path = os.environ.get("windir", "C:\\\\Windows") + "\\\\System32\\\\amdhip64_7.dll"
|
||||
if not os.path.isfile(path):
|
||||
path = os.environ.get("windir", "C:\\\\Windows") + "\\\\System32\\\\amdhip64_6.dll"
|
||||
if not os.path.isfile(path):
|
||||
path = os.environ.get("windir", "C:\\\\Windows") + "\\\\System32\\\\amdhip64.dll"
|
||||
assert os.path.isfile(path)
|
||||
self.handle = ctypes.windll.kernel32.LoadLibraryA(path.encode('utf-8'))
|
||||
ctypes.windll.kernel32.GetLastError.restype = ctypes.wintypes.DWORD
|
||||
ctypes.windll.kernel32.GetLastError.argtypes = []
|
||||
assert ctypes.windll.kernel32.GetLastError() == 0
|
||||
ctypes.windll.kernel32.GetProcAddress.restype = ctypes.c_void_p
|
||||
ctypes.windll.kernel32.GetProcAddress.argtypes = [ctypes.wintypes.HMODULE, ctypes.c_char_p]
|
||||
hipInit = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.c_uint)(
|
||||
ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipInit"))
|
||||
with mute(sys.stdout.fileno()):
|
||||
hipInit(0)
|
||||
self.hipGetDeviceCount = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(ctypes.c_int))(
|
||||
ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceCount"))
|
||||
self.hipGetDeviceProperties = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(hipDeviceProp), ctypes.c_int)(
|
||||
ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceProperties"))
|
||||
def __del__(self):
|
||||
if self.handle is None:
|
||||
return
|
||||
ctypes.windll.kernel32.FreeLibrary.argtypes = [ctypes.wintypes.HMODULE]
|
||||
ctypes.windll.kernel32.FreeLibrary(self.handle)
|
||||
def get_device_count(self) -> int:
|
||||
count = ctypes.c_int()
|
||||
assert self.hipGetDeviceCount(ctypes.byref(count)) == 0
|
||||
return count.value
|
||||
def get_device_properties(self, device_id) -> bytes:
|
||||
prop = hipDeviceProp()
|
||||
assert self.hipGetDeviceProperties(ctypes.byref(prop), device_id) == 0
|
||||
return bytes(prop)
|
||||
if __name__ == "__main__":
|
||||
if sys.platform != "win32":
|
||||
print("This script is only for Windows.")
|
||||
sys.exit(1)
|
||||
hip = HIP()
|
||||
count = hip.get_device_count()
|
||||
archs: list[str | None] = [None] * count
|
||||
for i in range(count):
|
||||
prop = hip.get_device_properties(i)
|
||||
name = ""
|
||||
idx = 0
|
||||
while idx < len(prop):
|
||||
try:
|
||||
idx = prop.index(0x67, idx) + 1
|
||||
except ValueError:
|
||||
break
|
||||
if prop[idx] != 0x66:
|
||||
continue
|
||||
if prop[idx + 1] != 0x78:
|
||||
continue
|
||||
idx = idx + 2
|
||||
while prop[idx] != 0x00:
|
||||
c = prop[idx]
|
||||
idx += 1
|
||||
if (c < 0x30 or c > 0x39) and (c < 0x61 or c > 0x66):
|
||||
name = ""
|
||||
continue
|
||||
name += chr(c)
|
||||
break
|
||||
if name:
|
||||
archs[i] = "gfx" + name
|
||||
del hip
|
||||
for arch in archs:
|
||||
if arch is not None:
|
||||
print(arch)
|
||||
"""
|
||||
|
||||
@@ -184,7 +184,7 @@ weights_dtype_order_fp32 = weights_dtype_order + [
|
||||
"uint16", "float16_e1m15fnu", "float16_e2m14fnu", "float16_e3m13fnu", "float16_e4m12fnu", "float16_e5m11fnu",
|
||||
]
|
||||
|
||||
is_rdna2 = bool(devices.backend == "rocm" and int(devices.get_hip_arch_name()) < 1100)
|
||||
is_rdna2 = bool(devices.backend == "rocm" and devices.get_hip_agent().gfx_version < 0x1100)
|
||||
use_torch_compile = shared.opts.sdnq_dequantize_compile # this setting requires a full restart of the webui to apply
|
||||
|
||||
def check_torch_compile(): # dynamo can be disabled after startup
|
||||
|
||||
@@ -52,7 +52,8 @@ def get_default_modes(cmd_opts, mem_stat):
|
||||
default_sdp_override_choices.append('Triton Flash attention')
|
||||
elif devices.backend == "rocm":
|
||||
default_sdp_override_choices.append('Triton Flash attention')
|
||||
if int(devices.get_hip_arch_name()) < 1100:
|
||||
agent = devices.get_hip_agent()
|
||||
if agent.gfx_version < 0x1100:
|
||||
default_sdp_override_options = ['Dynamic attention'] # only RDNA2 and older GPUs needs this
|
||||
elif devices.backend in {"directml", "cpu", "mps"}:
|
||||
default_sdp_override_options = ['Dynamic attention']
|
||||
|
||||
@@ -1,87 +0,0 @@
|
||||
import sys
|
||||
|
||||
if sys.platform == "win32":
|
||||
import os
|
||||
import ctypes
|
||||
import ctypes.wintypes
|
||||
|
||||
class hipDeviceProp(ctypes.Structure):
|
||||
_fields_ = [
|
||||
('bytes', ctypes.c_byte * 1472) # 1472 in amdhip64_6.dll, shorter in amdhip64_7.dll?
|
||||
]
|
||||
|
||||
class HIP:
|
||||
def __init__(self):
|
||||
ctypes.windll.kernel32.LoadLibraryA.restype = ctypes.wintypes.HMODULE
|
||||
ctypes.windll.kernel32.LoadLibraryA.argtypes = [ctypes.c_char_p]
|
||||
self.handle = None
|
||||
path = os.environ.get("windir", "C:\\Windows") + "\\System32\\amdhip64_7.dll"
|
||||
if not os.path.isfile(path):
|
||||
path = os.environ.get("windir", "C:\\Windows") + "\\System32\\amdhip64_6.dll"
|
||||
if not os.path.isfile(path):
|
||||
path = os.environ.get("windir", "C:\\Windows") + "\\System32\\amdhip64.dll"
|
||||
assert os.path.isfile(path)
|
||||
self.handle = ctypes.windll.kernel32.LoadLibraryA(path.encode('utf-8'))
|
||||
ctypes.windll.kernel32.GetLastError.restype = ctypes.wintypes.DWORD
|
||||
ctypes.windll.kernel32.GetLastError.argtypes = []
|
||||
assert ctypes.windll.kernel32.GetLastError() == 0
|
||||
ctypes.windll.kernel32.GetProcAddress.restype = ctypes.c_void_p
|
||||
ctypes.windll.kernel32.GetProcAddress.argtypes = [ctypes.wintypes.HMODULE, ctypes.c_char_p]
|
||||
self.hipGetDeviceCount = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(ctypes.c_int))(
|
||||
ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceCount"))
|
||||
self.hipGetDeviceProperties = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(hipDeviceProp), ctypes.c_int)(
|
||||
ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceProperties"))
|
||||
|
||||
def __del__(self):
|
||||
if self.handle is None:
|
||||
return
|
||||
# Hopefully this will prevent conflicts with amdhip64_7.dll from ROCm Python packages or HIP SDK
|
||||
ctypes.windll.kernel32.FreeLibrary.argtypes = [ctypes.wintypes.HMODULE]
|
||||
ctypes.windll.kernel32.FreeLibrary(self.handle)
|
||||
|
||||
def get_device_count(self):
|
||||
count = ctypes.c_int()
|
||||
assert self.hipGetDeviceCount(ctypes.byref(count)) == 0
|
||||
return count.value
|
||||
|
||||
def get_device_properties(self, device_id):
|
||||
prop = hipDeviceProp()
|
||||
assert self.hipGetDeviceProperties(ctypes.byref(prop), device_id) == 0
|
||||
return prop.bytes
|
||||
|
||||
def get_archs():
|
||||
hip = HIP()
|
||||
|
||||
count = hip.get_device_count()
|
||||
archs = [None] * count
|
||||
for i in range(count):
|
||||
prop = hip.get_device_properties(i)[:]
|
||||
|
||||
name = ""
|
||||
idx = 0
|
||||
while idx < len(prop):
|
||||
try:
|
||||
idx = prop.index(0x67, idx) + 1 # 'g'
|
||||
except ValueError:
|
||||
break
|
||||
if prop[idx] != 0x66: # 'f'
|
||||
continue
|
||||
if prop[idx + 1] != 0x78: # 'x'
|
||||
continue
|
||||
|
||||
idx = idx + 2
|
||||
while prop[idx] != 0x00:
|
||||
c = prop[idx]
|
||||
idx += 1
|
||||
if (c < 0x30 or c > 0x39) and (c < 0x61 or c > 0x66): # hexadecimal
|
||||
name = ""
|
||||
continue
|
||||
name += chr(c)
|
||||
break
|
||||
|
||||
# if name == "", hipDeviceProp does not contain arch name
|
||||
if name:
|
||||
archs[i] = "gfx" + name
|
||||
|
||||
del hip
|
||||
return archs
|
||||
Reference in New Issue
Block a user