diff --git a/unsloth/kernels/utils.py b/unsloth/kernels/utils.py index 23be372217..0af65edeee 100644 --- a/unsloth/kernels/utils.py +++ b/unsloth/kernels/utils.py @@ -54,6 +54,9 @@ pass import bitsandbytes as bnb +# https://github.com/bitsandbytes-foundation/bitsandbytes/pull/1330/files +HAS_CUDA_STREAM = Version(bnb.__version__) > Version("0.43.3") +CUDA_STREAM = torch.cuda.current_stream("cuda:0") get_ptr = bnb.functional.get_ptr import ctypes cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32 @@ -105,119 +108,237 @@ def get_lora_parameters_bias(proj): pass -def fast_dequantize(W, quant_state = None, out = None): - if quant_state is None: return W - 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 - 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 +if HAS_CUDA_STREAM: + def fast_dequantize(W, quant_state = None, out = None): + if quant_state is None: return W + 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 + 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 + pass + + # Create weight matrix + if out is None: + out = torch.empty(shape, dtype = dtype, device = "cuda:0") + else: + assert(out.shape == shape) + assert(out.dtype == dtype) + + # NF4 dequantization of statistics + n_elements_absmax = absmax.numel() + out_absmax = torch.empty(n_elements_absmax, dtype = torch.float32, device = "cuda:0") + + # 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), CUDA_STREAM, + ) + 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()), CUDA_STREAM,) + + # Careful returning transposed data + is_transposed = (True if W.shape[0] == 1 else False) + return out.t() if is_transposed else out pass +else: + def fast_dequantize(W, quant_state = None, out = None): + if quant_state is None: return W + 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 + 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 + pass - # Create weight matrix - if out is None: - out = torch.empty(shape, dtype = dtype, device = "cuda:0") - else: - assert(out.shape == shape) - assert(out.dtype == dtype) + # Create weight matrix + if out is None: + out = torch.empty(shape, dtype = dtype, device = "cuda:0") + else: + assert(out.shape == shape) + assert(out.dtype == dtype) - # NF4 dequantization of statistics - n_elements_absmax = absmax.numel() - out_absmax = torch.empty(n_elements_absmax, dtype = torch.float32, device = "cuda:0") + # NF4 dequantization of statistics + n_elements_absmax = absmax.numel() + out_absmax = torch.empty(n_elements_absmax, dtype = torch.float32, device = "cuda:0") - # 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 + # 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 - 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())) + 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()),) - # Careful returning transposed data - is_transposed = (True if W.shape[0] == 1 else False) - return out.t() if is_transposed else out + # Careful returning transposed data + is_transposed = (True if W.shape[0] == 1 else False) + return out.t() if is_transposed else out + pass pass -def fast_gemv(X, W, quant_state, out = None): - if quant_state is None: return torch.matmul(X, W, out = out) - # For fast X @ W where seq_len == 1 - # From https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L1469 - _, q_len, hd = X.shape - # assert(q_len == 1) +if HAS_CUDA_STREAM: + def fast_gemv(X, W, quant_state, out = None): + if quant_state is None: return torch.matmul(X, W, out = out) + # For fast X @ W where seq_len == 1 + # From https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L1469 + _, q_len, hd = X.shape + # assert(q_len == 1) - 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 - else: - absmax, shape, dtype, blocksize, compressed_stats, quant_type, stats = quant_state - offset, state2 = compressed_stats - absmax2, code2, blocksize2, _, _, _, _ = state2 + 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 + else: + absmax, shape, dtype, blocksize, compressed_stats, quant_type, stats = quant_state + offset, state2 = compressed_stats + absmax2, code2, blocksize2, _, _, _, _ = state2 + pass + # assert(dtype == X.dtype) + bout = shape[0] + + if out is None: + out = torch.empty((1, 1, bout,), dtype = dtype, device = "cuda:0") + # else: + # assert(out.shape == (1, 1, bout,)) + # pass + + n = 1 + m = shape[0] + k = shape[1] + lda = shape[0] + ldc = shape[0] + ldb = (hd+1)//2 + m = ctypes.c_int32(m) + n = ctypes.c_int32(n) + k = ctypes.c_int32(k) + lda = ctypes.c_int32(lda) + ldb = ctypes.c_int32(ldb) + ldc = ctypes.c_int32(ldc) + + df = torch.empty(absmax.shape, dtype = torch.float32, device = "cuda:0") + 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 + + blocksize = ctypes.c_int32(blocksize) + fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out), + lda, ldb, ldc, blocksize, CUDA_STREAM,) + + return out pass - # assert(dtype == X.dtype) - bout = shape[0] +else: + def fast_gemv(X, W, quant_state, out = None): + if quant_state is None: return torch.matmul(X, W, out = out) + # For fast X @ W where seq_len == 1 + # From https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L1469 + _, q_len, hd = X.shape + # assert(q_len == 1) - if out is None: - out = torch.empty((1, 1, bout,), dtype = dtype, device = "cuda:0") - # else: - # assert(out.shape == (1, 1, bout,)) - # pass + 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 + else: + absmax, shape, dtype, blocksize, compressed_stats, quant_type, stats = quant_state + offset, state2 = compressed_stats + absmax2, code2, blocksize2, _, _, _, _ = state2 + pass + # assert(dtype == X.dtype) + bout = shape[0] - n = 1 - m = shape[0] - k = shape[1] - lda = shape[0] - ldc = shape[0] - ldb = (hd+1)//2 - m = ctypes.c_int32(m) - n = ctypes.c_int32(n) - k = ctypes.c_int32(k) - lda = ctypes.c_int32(lda) - ldb = ctypes.c_int32(ldb) - ldc = ctypes.c_int32(ldc) + if out is None: + out = torch.empty((1, 1, bout,), dtype = dtype, device = "cuda:0") + # else: + # assert(out.shape == (1, 1, bout,)) + # pass - df = torch.empty(absmax.shape, dtype = torch.float32, device = "cuda:0") - 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 + n = 1 + m = shape[0] + k = shape[1] + lda = shape[0] + ldc = shape[0] + ldb = (hd+1)//2 + m = ctypes.c_int32(m) + n = ctypes.c_int32(n) + k = ctypes.c_int32(k) + lda = ctypes.c_int32(lda) + ldb = ctypes.c_int32(ldb) + ldc = ctypes.c_int32(ldc) - fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \ - cgemm_4bit_inference_naive_bf16 + df = torch.empty(absmax.shape, dtype = torch.float32, device = "cuda:0") + 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 - blocksize = ctypes.c_int32(blocksize) - fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out), - lda, ldb, ldc, blocksize) + fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \ + cgemm_4bit_inference_naive_bf16 - return out + blocksize = ctypes.c_int32(blocksize) + fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out), + lda, ldb, ldc, blocksize,) + + return out + pass pass diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index ed6207bb06..15e9efc426 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -181,6 +181,7 @@ def MistralForCausalLM_fast_forward( output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, + num_logits_to_keep: Optional[int] = 0, *args, **kwargs, ) -> Union[Tuple, CausalLMOutputWithPast]: @@ -236,6 +237,8 @@ def MistralForCausalLM_fast_forward( if bsz == 1 and q_len == 1: logits = torch.mv(lm_head, hidden_states.ravel().to(lm_head.dtype)) logits = logits.unsqueeze(0).unsqueeze(0) + elif num_logits_to_keep != 0: + logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :].to(lm_head.dtype)) else: logits = self.lm_head(hidden_states.to(lm_head.dtype)) pass