diff --git a/unsloth/kernels/utils.py b/unsloth/kernels/utils.py index 645319d423..eac2b59974 100644 --- a/unsloth/kernels/utils.py +++ b/unsloth/kernels/utils.py @@ -61,7 +61,6 @@ else: pass pass - def calculate_settings(n : int) -> (int, int,): BLOCK_SIZE : int = next_power_of_2(n) if BLOCK_SIZE > MAX_FUSED_SIZE: @@ -87,7 +86,7 @@ else: # https://github.com/bitsandbytes-foundation/bitsandbytes/pull/1330/files HAS_CUDA_STREAM = Version(bnb.__version__) > Version("0.43.3") get_ptr = bnb.functional.get_ptr - +pass if DEVICE_COUNT > 1: if DEVICE_TYPE == "cuda": @@ -97,7 +96,7 @@ if DEVICE_COUNT > 1: else: from contextlib import nullcontext def torch_gpu_device(device): return nullcontext() - pass +pass # INTEL GPU Specific Logic if DEVICE_TYPE == "xpu": @@ -105,13 +104,13 @@ if DEVICE_TYPE == "xpu": # NVIDIA GPU Default Logic else: _gpu_getCurrentRawStream = torch._C._cuda_getCurrentRawStream +pass c_void_p = ctypes.c_void_p def _get_tensor_stream(tensor: torch_Tensor) -> c_void_p: return c_void_p(_gpu_getCurrentRawStream(tensor.device.index)) pass - # Get array of CUDA streams and other buffers global CUDA_STREAMS global XPU_STREAMS @@ -124,7 +123,7 @@ if DEVICE_TYPE == "xpu": (index := torch.xpu.device(i).idx) : ctypes.c_void_p(torch._C._xpu_getCurrentRawStream(index)) for i in range(DEVICE_COUNT) } - XPU_STREAMS = [None] * (max(_XPU_STREAMS.keys()) + 1) + XPU_STREAMS = [None] * (max(_XPU_STREAMS.keys()) + 1) WEIGHT_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1) ABSMAX_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1) for k, v in _XPU_STREAMS.items(): @@ -143,7 +142,7 @@ else: for k, v in _CUDA_STREAMS.items(): CUDA_STREAMS[k] = v CUDA_STREAMS = tuple(CUDA_STREAMS) del _CUDA_STREAMS - +pass # Bitsandbytes operations ctypes_c_int = ctypes.c_int @@ -172,12 +171,15 @@ else: cdequantize_blockwise_bf16_nf4 = bnb.functional.lib.cdequantize_blockwise_bf16_nf4 cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemm_4bit_inference_naive_fp16 cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemm_4bit_inference_naive_bf16 +pass torch_mm = torch.mm torch_mv = torch.mv -torch_matmul = torch.matmul -torch_addmm = torch.addmm -torch_empty = torch.empty +torch_matmul = torch.matmul +torch_addmm = torch.addmm +torch_empty = torch.empty +torch_float16 = torch.float16 +torch_float32 = torch.float32 def QUANT_STATE(W): return getattr(W, "quant_state", None) @@ -235,23 +237,28 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM: def fast_dequantize(W, quant_state = None, out = None, use_global_buffer = False): # TODO: After adding XPU BNB support, check this function if quant_state is None: return W + is_double_quantized = True if type(quant_state) is not list: # New quant_state as a class # https://github.com/TimDettmers/bitsandbytes/pull/763/files - absmax = quant_state.absmax - shape = quant_state.shape - dtype = quant_state.dtype - blocksize = quant_state.blocksize - offset = quant_state.offset - state2 = quant_state.state2 - absmax2 = state2.absmax - code2 = state2.code - blocksize2 = state2.blocksize + absmax = quant_state.absmax + shape = quant_state.shape + dtype = quant_state.dtype + blocksize = quant_state.blocksize + offset = quant_state.offset + state2 = quant_state.state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2 = state2.absmax + code2 = state2.code + blocksize2 = state2.blocksize else: # Old quant_state as a list of lists absmax, shape, dtype, blocksize, compressed_stats, _, _ = quant_state offset, state2 = compressed_stats - absmax2, code2, blocksize2, _, _, _, _ = state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2, code2, blocksize2, _, _, _, _ = state2 pass global XPU_STREAMS device = W.device @@ -270,7 +277,7 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM: ABSMAX_BUFFER = ABSMAX_BUFFERS[device_index] if WEIGHT_BUFFER is None: WEIGHT_BUFFERS[device_index] = WEIGHT_BUFFER = torch_empty(size, dtype = dtype, device = device, requires_grad = False) - ABSMAX_BUFFERS[device_index] = ABSMAX_BUFFER = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False) + ABSMAX_BUFFERS[device_index] = ABSMAX_BUFFER = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False) if size > WEIGHT_BUFFER.numel(): WEIGHT_BUFFER.resize_(size) if n_elements_absmax > ABSMAX_BUFFER.numel(): ABSMAX_BUFFER.resize_(n_elements_absmax) @@ -283,20 +290,23 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM: else: assert(out.shape == shape) assert(out.dtype == dtype) - out_absmax = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False) + out_absmax = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False) pass # NF4 dequantization of statistics - ptr_out_absmax = get_ptr(out_absmax) with torch_gpu_device(device): - 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), XPU_STREAM - ) - out_absmax += offset + if is_double_quantized: + 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), XPU_STREAM + ) + out_absmax += offset + else: + ptr_out_absmax = get_ptr(absmax) # Dequantize W - fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \ + 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()), XPU_STREAM,) @@ -310,23 +320,28 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM: @torch.inference_mode def fast_dequantize(W, quant_state = None, out = None, use_global_buffer = False): if quant_state is None: return W + is_double_quantized = True if type(quant_state) is not list: # New quant_state as a class # https://github.com/TimDettmers/bitsandbytes/pull/763/files - absmax = quant_state.absmax - shape = quant_state.shape - dtype = quant_state.dtype - blocksize = quant_state.blocksize - offset = quant_state.offset - state2 = quant_state.state2 - absmax2 = state2.absmax - code2 = state2.code - blocksize2 = state2.blocksize + absmax = quant_state.absmax + shape = quant_state.shape + dtype = quant_state.dtype + blocksize = quant_state.blocksize + offset = quant_state.offset + state2 = quant_state.state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2 = state2.absmax + code2 = state2.code + blocksize2 = state2.blocksize else: # Old quant_state as a list of lists absmax, shape, dtype, blocksize, compressed_stats, _, _ = quant_state offset, state2 = compressed_stats - absmax2, code2, blocksize2, _, _, _, _ = state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2, code2, blocksize2, _, _, _, _ = state2 pass global CUDA_STREAMS device = W.device @@ -346,7 +361,7 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM: ABSMAX_BUFFER = ABSMAX_BUFFERS[device_index] if WEIGHT_BUFFER is None: WEIGHT_BUFFERS[device_index] = WEIGHT_BUFFER = torch_empty(size, dtype = dtype, device = device, requires_grad = False) - ABSMAX_BUFFERS[device_index] = ABSMAX_BUFFER = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False) + ABSMAX_BUFFERS[device_index] = ABSMAX_BUFFER = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False) if size > WEIGHT_BUFFER.numel(): WEIGHT_BUFFER.resize_(size) if n_elements_absmax > ABSMAX_BUFFER.numel(): ABSMAX_BUFFER.resize_(n_elements_absmax) @@ -359,20 +374,22 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM: else: assert(out.shape == shape) assert(out.dtype == dtype) - out_absmax = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False) + out_absmax = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False) pass # NF4 dequantization of statistics - ptr_out_absmax = get_ptr(out_absmax) with torch_gpu_device(device): - 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), CUDA_STREAM - ) - out_absmax += offset - + if is_double_quantized: + 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), CUDA_STREAM + ) + out_absmax += offset + else: + ptr_out_absmax = get_ptr(absmax) # Dequantize W - fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \ + 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()), CUDA_STREAM,) @@ -385,23 +402,28 @@ else: @torch.inference_mode def fast_dequantize(W, quant_state = None, out = None, use_global_buffer = False): if quant_state is None: return W + is_double_quantized = True if type(quant_state) is not list: # New quant_state as a class # https://github.com/TimDettmers/bitsandbytes/pull/763/files - absmax = quant_state.absmax - shape = quant_state.shape - dtype = quant_state.dtype - blocksize = quant_state.blocksize - offset = quant_state.offset - state2 = quant_state.state2 - absmax2 = state2.absmax - code2 = state2.code - blocksize2 = state2.blocksize + absmax = quant_state.absmax + shape = quant_state.shape + dtype = quant_state.dtype + blocksize = quant_state.blocksize + offset = quant_state.offset + state2 = quant_state.state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2 = state2.absmax + code2 = state2.code + blocksize2 = state2.blocksize else: # Old quant_state as a list of lists absmax, shape, dtype, blocksize, compressed_stats, _, _ = quant_state offset, state2 = compressed_stats - absmax2, code2, blocksize2, _, _, _, _ = state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2, code2, blocksize2, _, _, _, _ = state2 pass n_elements_absmax = absmax.numel() @@ -413,17 +435,20 @@ else: else: assert(out.shape == shape) assert(out.dtype == dtype) - out_absmax = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False) + out_absmax = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False) # Do dequantization - 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), - ) - out_absmax += offset + if is_double_quantized: + 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), + ) + out_absmax += offset + else: + ptr_out_absmax = get_ptr(absmax) - fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \ + 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()),) @@ -443,23 +468,27 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM: # From https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L1469 _, q_len, hd = X.shape # assert(q_len == 1) - + is_double_quantized = True if type(quant_state) is not list: # https://github.com/TimDettmers/bitsandbytes/pull/763/files - absmax = quant_state.absmax - shape = quant_state.shape - dtype = quant_state.dtype - blocksize = quant_state.blocksize - stats = quant_state.code - offset = quant_state.offset - state2 = quant_state.state2 - absmax2 = state2.absmax - code2 = state2.code - blocksize2 = state2.blocksize + absmax = quant_state.absmax + shape = quant_state.shape + dtype = quant_state.dtype + blocksize = quant_state.blocksize + stats = quant_state.code + offset = quant_state.offset + state2 = quant_state.state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2 = state2.absmax + code2 = state2.code + blocksize2 = state2.blocksize else: absmax, shape, dtype, blocksize, compressed_stats, quant_type, stats = quant_state offset, state2 = compressed_stats - absmax2, code2, blocksize2, _, _, _, _ = state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2, code2, blocksize2, _, _, _, _ = state2 pass global XPU_STREAMS device = W.device @@ -488,17 +517,18 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM: ldb = ctypes_c_int32(ldb) ldc = ctypes_c_int32(ldc) - df = torch_empty(absmax.shape, dtype = torch.float32, device = device) with torch_gpu_device(device): - cdequantize_blockwise_fp32( - get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df), - ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), XPU_STREAM, - ) - df += offset - absmax = df + if is_double_quantized: + df = torch_empty(absmax.shape, dtype = torch_float32, device = device) + cdequantize_blockwise_fp32( + get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df), + ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), XPU_STREAM, + ) + df += offset + absmax = df - fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \ - cgemm_4bit_inference_naive_bf16 + fx = cgemm_4bit_inference_naive_fp16 if dtype == torch_float16 else \ + cgemm_4bit_inference_naive_bf16 blocksize = ctypes_c_int32(blocksize) fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out), @@ -514,23 +544,28 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM: # From https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L1469 _, q_len, hd = X.shape # assert(q_len == 1) + is_double_quantized = True if type(quant_state) is not list: # https://github.com/TimDettmers/bitsandbytes/pull/763/files - absmax = quant_state.absmax - shape = quant_state.shape - dtype = quant_state.dtype - blocksize = quant_state.blocksize - stats = quant_state.code - offset = quant_state.offset - state2 = quant_state.state2 - absmax2 = state2.absmax - code2 = state2.code - blocksize2 = state2.blocksize + absmax = quant_state.absmax + shape = quant_state.shape + dtype = quant_state.dtype + blocksize = quant_state.blocksize + stats = quant_state.code + offset = quant_state.offset + state2 = quant_state.state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2 = state2.absmax + code2 = state2.code + blocksize2 = state2.blocksize else: absmax, shape, dtype, blocksize, compressed_stats, quant_type, stats = quant_state offset, state2 = compressed_stats - absmax2, code2, blocksize2, _, _, _, _ = state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2, code2, blocksize2, _, _, _, _ = state2 pass global CUDA_STREAMS device = W.device @@ -559,17 +594,18 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM: ldb = ctypes_c_int32(ldb) ldc = ctypes_c_int32(ldc) - df = torch_empty(absmax.shape, dtype = torch.float32, device = device) with torch_gpu_device(device): - cdequantize_blockwise_fp32( - get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df), - ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), CUDA_STREAM, - ) - df += offset - absmax = df + if is_double_quantized: + df = torch_empty(absmax.shape, dtype = torch_float32, device = device) + cdequantize_blockwise_fp32( + get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df), + ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), CUDA_STREAM, + ) + df += offset + absmax = df - fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \ - cgemm_4bit_inference_naive_bf16 + fx = cgemm_4bit_inference_naive_fp16 if dtype == torch_float16 else \ + cgemm_4bit_inference_naive_bf16 blocksize = ctypes_c_int32(blocksize) fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out), @@ -585,6 +621,7 @@ else: # From https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L1469 _, q_len, hd = X.shape # assert(q_len == 1) + is_double_quantized = True if type(quant_state) is not list: # https://github.com/TimDettmers/bitsandbytes/pull/763/files @@ -595,13 +632,17 @@ else: stats = quant_state.code offset = quant_state.offset state2 = quant_state.state2 - absmax2 = state2.absmax - code2 = state2.code - blocksize2 = state2.blocksize + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2 = state2.absmax + code2 = state2.code + blocksize2 = state2.blocksize else: absmax, shape, dtype, blocksize, compressed_stats, quant_type, stats = quant_state offset, state2 = compressed_stats - absmax2, code2, blocksize2, _, _, _, _ = state2 + is_double_quantized = state2 is not None + if is_double_quantized: + absmax2, code2, blocksize2, _, _, _, _ = state2 pass # assert(dtype == X.dtype) bout = shape[0] @@ -626,16 +667,17 @@ else: ldb = ctypes_c_int32(ldb) ldc = ctypes_c_int32(ldc) - df = torch_empty(absmax.shape, dtype = torch.float32, device = device) - cdequantize_blockwise_fp32( - get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df), - ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), - ) - df += offset - absmax = df + if is_double_quantized: + df = torch_empty(absmax.shape, dtype = torch_float32, device = device) + cdequantize_blockwise_fp32( + get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df), + ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), + ) + df += offset + absmax = df - fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \ - cgemm_4bit_inference_naive_bf16 + fx = cgemm_4bit_inference_naive_fp16 if dtype == torch_float16 else \ + cgemm_4bit_inference_naive_bf16 blocksize = ctypes_c_int32(blocksize) fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out),