From 67819bb57cc7c72b58d5284af870aea45282b1c0 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 2 Feb 2025 17:07:31 -0800 Subject: [PATCH] Update utils.py --- unsloth/kernels/utils.py | 19 ++++++++++++------- 1 file changed, 12 insertions(+), 7 deletions(-) diff --git a/unsloth/kernels/utils.py b/unsloth/kernels/utils.py index 8537e9595e..3b0c1d3919 100644 --- a/unsloth/kernels/utils.py +++ b/unsloth/kernels/utils.py @@ -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)