From ad02ea81b75a2bdb11ad3133ca89226d8a9ac0e2 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sat, 30 Dec 2023 04:35:35 +1100 Subject: [PATCH] Prelim Qwen, Deepseek support (#58) * Pytorch 2.1.1 install path, 4bit loading * Update README.md * Update README.md * Update README.md * Update README.md * Update README.md * Update README.md * Update README.md * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Spelling errors * Update __init__.py * DPO loss fix * Update fast_lora.py * Update fast_lora.py * Out of bounds tokenization * Fix Mistral SWA * Prelim support Qwen, Deepseek etc --- unsloth/kernels/cross_entropy_loss.py | 32 ++++++++++++++----- unsloth/kernels/fast_lora.py | 7 +++-- unsloth/models/_utils.py | 2 +- unsloth/models/llama.py | 44 ++++++++++++++++++++++----- 4 files changed, 67 insertions(+), 18 deletions(-) diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index 67ad306fcf..0a73a393ec 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -15,7 +15,9 @@ import triton import triton.language as tl import torch -from .utils import calculate_settings +from .utils import calculate_settings, MAX_FUSED_SIZE +from transformers.models.llama.modeling_llama import logger + @triton.jit def _cross_entropy_forward(logits_ptr, logits_row_stride, @@ -145,6 +147,7 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): pass +slow_cross_entropy_loss = torch.nn.functional.cross_entropy def fast_cross_entropy_loss(logits, labels): """ Arguments: @@ -156,10 +159,25 @@ def fast_cross_entropy_loss(logits, labels): batch, seq_len, d = logits.shape assert(labels.shape == (batch, seq_len)) - loss = Fast_CrossEntropyLoss.apply( - logits.view(batch*seq_len, d), - labels.view(-1), - ) - n_items = torch.count_nonzero(labels != -100) - return loss.sum() / n_items + # Prelim support Qwen, Deepseek other large vocab sizes > 2^16 + if d > MAX_FUSED_SIZE: + logger.warning_once( + f"Unsloth: Vocab size of {d} exceeds the max CUDA blocksize of {MAX_FUSED_SIZE}.\n"\ + "For now, Unsloth will use Pytorch's CrossEntropyLoss, which will entail a\n"\ + "25% increase in memory usage and be slower. Make an issue on \n"\ + "Unsloth's Github page if you want a faster and more memory efficient kernel!" + ) + loss = slow_cross_entropy_loss( + logits.float().view(batch*seq_len, d), # Must cast to float32 for numerical stability + labels.view(-1), + ) + return loss + else: + loss = Fast_CrossEntropyLoss.apply( + logits.view(batch*seq_len, d), + labels.view(-1), + ) + n_items = torch.count_nonzero(labels != -100) + return loss.sum() / n_items + pass pass diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index adec68d46b..26fecf8b59 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -16,6 +16,7 @@ import torch from .utils import fast_dequantize, QUANT_STATE from .swiglu import swiglu_fg_kernel, swiglu_DWf_DW_dfg_kernel + def get_lora_parameters(proj): # For DPO or disabled adapters base_layer = (proj.base_layer if hasattr(proj, "base_layer") else proj) @@ -104,7 +105,7 @@ class LoRA_MLP(torch.autograd.Function): dtype = X.dtype e = matmul_lora(X, gateW, gateW_quant, gateA, gateB, gateS) - g = matmul_lora(X, upW, upW_quant, upA, upB, upS) + g = matmul_lora(X, upW, upW_quant, upA, upB, upS) h = swiglu_fg_kernel(e, g) i = matmul_lora(h, downW, downW_quant, downA, downB, downS) @@ -123,10 +124,10 @@ class LoRA_MLP(torch.autograd.Function): def backward(ctx, dY : torch.Tensor): gateW, gateW_quant, gateS, upW, upW_quant, upS, downW, downW_quant, downS, = \ ctx.custom_saved_tensors - gateA, gateB, upA,upB, downA, downB, \ + gateA, gateB, upA, upB, downA, downB, \ X, e, g = ctx.saved_tensors - gateA, gateB, upA,upB, downA, downB = \ + gateA, gateB, upA, upB, downA, downB = \ gateA.t(), gateB.t(), upA.t(), upB.t(), downA.t(), downB.t() batch, seq_len, hd = X.shape diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 769669ae6a..1c75dffd53 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -23,7 +23,7 @@ from transformers.models.llama.modeling_llama import logger from platform import system as platform_system platform_system = platform_system() -__version__ = "2023.12" +__version__ = "2024.1" # Get Flash Attention v2 if Ampere (RTX 30xx, A100) major_version, minor_version = torch.cuda.get_device_capability() diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 22179d4410..ccf61ee920 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -817,29 +817,59 @@ class FastLlamaModel: for idx, layer in enumerate(model.model.model.layers): # MLP patching - if hasattr(layer.mlp.gate_proj, "lora_A") and \ - hasattr(layer.mlp. up_proj, "lora_A") and \ - hasattr(layer.mlp.down_proj, "lora_A"): + gate_proj = layer.mlp.gate_proj + up_proj = layer.mlp. up_proj + down_proj = layer.mlp.down_proj + + if hasattr(gate_proj, "lora_A") and \ + hasattr( up_proj, "lora_A") and \ + hasattr(down_proj, "lora_A") and \ + (gate_proj.base_layer if hasattr(gate_proj, "base_layer") else gate_proj).bias is None and \ + ( up_proj.base_layer if hasattr( up_proj, "base_layer") else up_proj).bias is None and \ + (down_proj.base_layer if hasattr(down_proj, "base_layer") else down_proj).bias is None: # https://stackoverflow.com/questions/50599045/python-replacing-a-function-within-a-class-of-a-module layer.mlp.forward = types.MethodType(apply_lora_mlp, layer.mlp) n_mlp += 1 + else: + logger.warning_once( + "Unsloth cannot patch MLP layers with our manual autograd engine since either LoRA adapters\n"\ + "are not enabled or a bias term (like in Qwen) is used." + ) pass # QKV attention patching - if hasattr(layer.self_attn.q_proj, "lora_A") and \ - hasattr(layer.self_attn.k_proj, "lora_A") and \ - hasattr(layer.self_attn.v_proj, "lora_A"): + q_proj = layer.self_attn.q_proj + k_proj = layer.self_attn.k_proj + v_proj = layer.self_attn.v_proj + if hasattr(q_proj, "lora_A") and \ + hasattr(k_proj, "lora_A") and \ + hasattr(v_proj, "lora_A") and \ + (q_proj.base_layer if hasattr(q_proj, "base_layer") else q_proj).bias is None and \ + (k_proj.base_layer if hasattr(k_proj, "base_layer") else k_proj).bias is None and \ + (v_proj.base_layer if hasattr(v_proj, "base_layer") else v_proj).bias is None: layer.self_attn.apply_qkv = apply_lora_qkv n_qkv += 1 + else: + logger.warning_once( + "Unsloth cannot patch Attention layers with our manual autograd engine since either LoRA adapters\n"\ + "are not enabled or a bias term (like in Qwen) is used." + ) pass # O attention patching - if hasattr(layer.self_attn.o_proj, "lora_A"): + o_proj = layer.self_attn.o_proj + if hasattr(o_proj, "lora_A") and \ + (o_proj.base_layer if hasattr(o_proj, "base_layer") else o_proj).bias is None: layer.self_attn.apply_o = apply_lora_o n_o += 1 + else: + logger.warning_once( + "Unsloth cannot patch O projection layer with our manual autograd engine since either LoRA adapters\n"\ + "are not enabled or a bias term (like in Qwen) is used." + ) pass pass