hip multidevice (#1728)

* feat: hip multidevice support + p2p

* feat: default device
This commit is contained in:
wozeparrot
2023-09-01 06:46:13 -07:00
committed by GitHub
parent 024dd690fa
commit bf05534c6e
2 changed files with 35 additions and 10 deletions
+13
View File
@@ -115,6 +115,9 @@ hipMemcpyDefault = 4
_libhip.hipMemcpy.restype = int
_libhip.hipMemcpy.argtypes = [ctypes.c_void_p, ctypes.c_void_p, ctypes.c_size_t, ctypes.c_int]
def hipMemcpy(dst, src, count, direction):
status = _libhip.hipMemcpy(dst, src, ctypes.c_size_t(count), direction)
hipCheckStatus(status)
def hipMemcpy_htod(dst, src, count):
status = _libhip.hipMemcpy(dst, src, ctypes.c_size_t(count), hipMemcpyHostToDevice)
@@ -178,6 +181,16 @@ def hipGetDevice():
return dev.value
_libhip.hipGetDeviceCount.restype = int
_libhip.hipGetDeviceCount.argtypes = [ctypes.POINTER(ctypes.c_int)]
def hipGetDeviceCount():
count = ctypes.c_int()
status = _libhip.hipGetDeviceCount(ctypes.byref(count))
hipCheckStatus(status)
return count.value
class hipDeviceArch(ctypes.Structure):
_fields_ = [
# *32-bit Atomics*
+22 -10
View File
@@ -1,9 +1,9 @@
import numpy as np
import ctypes, functools
import extra.hip_wrapper as hip
from tinygrad.helpers import DEBUG
from tinygrad.helpers import DEBUG, getenv
from tinygrad.ops import Compiled
from tinygrad.runtime.lib import RawBufferCopyInOut, LRUAllocator
from tinygrad.runtime.lib import RawBufferCopyInOut, LRUAllocator, RawBufferTransfer
from tinygrad.codegen.kernel import LinearizerOptions
from tinygrad.renderer.cstyle import uops_to_cstyle, CStyleLanguage
@@ -15,22 +15,31 @@ if DEBUG >= 5:
# The default HIP stream is used for everything.
class HIPAllocator(LRUAllocator):
def _do_alloc(self, size, dtype, device, **kwargs): return hip.hipMalloc(size * dtype.itemsize)
def _do_alloc(self, size, dtype, device, **kwargs):
hip.hipSetDevice(device)
return hip.hipMalloc(size * dtype.itemsize)
def _do_free(self, buf): hip.hipFree(buf)
def _cached_bufkey(self, size, dtype, device): return (device, size*dtype.itemsize) # Buffers of the same length could be reused, no matter what dtype.
HIPAlloc = HIPAllocator(hip.hipGetDeviceProperties(hip.hipGetDevice()).totalGlobalMem)
class RawHIPBuffer(RawBufferCopyInOut):
def __init__(self, size, dtype): super().__init__(size, dtype, allocator=HIPAlloc)
class _HIP:
def __init__(self):
self.device_count = hip.hipGetDeviceCount()
self.default_device = getenv("HIP_DEFAULT_DEVICE")
self.allocator = HIPAllocator(hip.hipGetDeviceProperties(self.default_device).totalGlobalMem)
HIP = _HIP()
class RawHIPBuffer(RawBufferCopyInOut, RawBufferTransfer):
def __init__(self, size, dtype, device=str(HIP.default_device)): super().__init__(size, dtype, allocator=HIP.allocator, **{'device': int(device)})
def _copyin(self, x:np.ndarray): hip.hipMemcpyAsync_htod(self._buf, x.ctypes.data, self.size * self.dtype.itemsize, 0)
def _copyout(self, x:np.ndarray): hip.hipMemcpy_dtoh(x.ctypes.data, self._buf, self.size * self.dtype.itemsize)
def _transfer(self, x): hip.hipMemcpyAsync(self._buf, x._buf, self.size * self.dtype.itemsize, hip.hipMemcpyDeviceToDevice, 0)
class HIPProgram:
def __init__(self, name:str, prg:str, binary=False):
try:
if not binary:
prog = hip.hiprtcCreateProgram(prg, name, [], [])
device_properties = hip.hipGetDeviceProperties(hip.hipGetDevice())
device_properties = hip.hipGetDeviceProperties(HIP.default_device)
hip.hiprtcCompileProgram(prog, [f'--offload-arch={device_properties.gcnArchName}'])
prg = hip.hiprtcGetCode(prog)
except Exception as e:
@@ -40,17 +49,20 @@ class HIPProgram:
asm = early_exec((["/opt/rocm/llvm/bin/llvm-objdump", '-d', '-'], prg))
print('\n'.join([x for x in asm.decode('utf-8').split("\n") if 's_code_end' not in x]))
module = hip.hipModuleLoadData(prg)
self.prg = hip.hipModuleGetFunction(module, name)
self.prgs = []
for i in range(HIP.device_count):
hip.hipSetDevice(i)
self.prgs.append(hip.hipModuleGetFunction(hip.hipModuleLoadData(prg), name))
def __call__(self, global_size, local_size, *args, wait=False):
hip.hipSetDevice(args[0]._device)
if wait:
start, end = hip.hipEventCreate(), hip.hipEventCreate()
hip.hipEventRecord(start)
class PackageStruct(ctypes.Structure):
_fields_ = [(f'field{idx}', ctypes.c_void_p) for idx in range(len(args))]
struct = PackageStruct(*[data._buf for data in args])
hip.hipModuleLaunchKernel(self.prg, global_size[0], global_size[1], global_size[2], local_size[0], local_size[1], local_size[2], 0, 0, struct)
hip.hipModuleLaunchKernel(self.prgs[args[0]._device], global_size[0], global_size[1], global_size[2], local_size[0], local_size[1], local_size[2], 0, 0, struct)
if wait:
hip.hipEventRecord(end)
hip.hipEventSynchronize(end)