Update utils.py
This commit is contained in:
parent
19ae4b1864
commit
67819bb57c
1 changed files with 12 additions and 7 deletions
|
|
@ -121,6 +121,8 @@ WEIGHT_BUFFER = None
|
|||
global ABSMAX_BUFFER
|
||||
ABSMAX_BUFFER = None
|
||||
|
||||
ctypes_c_int = ctypes.c_int
|
||||
|
||||
if HAS_CUDA_STREAM:
|
||||
def fast_dequantize(W, quant_state = None, out = None, use_global_buffer = False):
|
||||
if quant_state is None: return W
|
||||
|
|
@ -157,12 +159,14 @@ if HAS_CUDA_STREAM:
|
|||
if WEIGHT_BUFFER is None:
|
||||
WEIGHT_BUFFER = torch.empty(size, dtype = dtype, device = "cuda:0")
|
||||
ABSMAX_BUFFER = torch.empty(n_elements_absmax, dtype = torch.float32, device = "cuda:0")
|
||||
ABSMAX_BUFFER.ptr_out_absmax = get_ptr(ABSMAX_BUFFER)
|
||||
|
||||
if size > WEIGHT_BUFFER.numel(): WEIGHT_BUFFER.resize_(size)
|
||||
if n_elements_absmax > ABSMAX_BUFFER.numel(): ABSMAX_BUFFER.resize_(n_elements_absmax)
|
||||
|
||||
out = WEIGHT_BUFFER[:size].view(shape)
|
||||
out_absmax = ABSMAX_BUFFER[:n_elements_absmax]
|
||||
ptr_out_absmax = ABSMAX_BUFFER.ptr_out_absmax
|
||||
else:
|
||||
if out is None:
|
||||
out = torch.empty(shape, dtype = dtype, device = "cuda:0")
|
||||
|
|
@ -170,19 +174,20 @@ if HAS_CUDA_STREAM:
|
|||
assert(out.shape == shape)
|
||||
assert(out.dtype == dtype)
|
||||
out_absmax = torch.empty(n_elements_absmax, dtype = torch.float32, device = "cuda:0")
|
||||
ptr_out_absmax = get_ptr(out_absmax)
|
||||
pass
|
||||
|
||||
# NF4 dequantization of statistics
|
||||
ptr_out_absmax = get_ptr(out_absmax)
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2),
|
||||
get_ptr(absmax),
|
||||
get_ptr(absmax2),
|
||||
ptr_out_absmax,
|
||||
ctypes.c_int(blocksize2),
|
||||
ctypes.c_int(n_elements_absmax),
|
||||
ctypes_c_int(blocksize2),
|
||||
ctypes_c_int(n_elements_absmax),
|
||||
CUDA_STREAM,
|
||||
)
|
||||
print(offset, out_absmax)
|
||||
out_absmax += offset
|
||||
|
||||
# Dequantize W
|
||||
|
|
@ -193,8 +198,8 @@ if HAS_CUDA_STREAM:
|
|||
get_ptr(W),
|
||||
ptr_out_absmax,
|
||||
get_ptr(out),
|
||||
ctypes.c_int(blocksize),
|
||||
ctypes.c_int(out.numel()),
|
||||
ctypes_c_int(blocksize),
|
||||
ctypes_c_int(out.numel()),
|
||||
CUDA_STREAM,)
|
||||
|
||||
# Careful returning transposed data
|
||||
|
|
@ -254,14 +259,14 @@ else:
|
|||
ptr_out_absmax = get_ptr(out_absmax)
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), ptr_out_absmax,
|
||||
ctypes.c_int(blocksize2), ctypes.c_int(n_elements_absmax),
|
||||
ctypes_c_int(blocksize2), ctypes_c_int(n_elements_absmax),
|
||||
)
|
||||
out_absmax += offset
|
||||
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \
|
||||
cdequantize_blockwise_bf16_nf4
|
||||
fx(get_ptr(None), get_ptr(W), ptr_out_absmax, get_ptr(out),
|
||||
ctypes.c_int(blocksize), ctypes.c_int(out.numel()),)
|
||||
ctypes_c_int(blocksize), ctypes_c_int(out.numel()),)
|
||||
|
||||
# Careful returning transposed data
|
||||
is_transposed = (True if W.shape[0] == 1 else False)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue