diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index 3ed0d3c914..6568bba681 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -183,8 +183,8 @@ def apply_lora_mlp_swiglu(self, X): pass -from .geglu import geglu_forward_kernel, geglu_backward_kernel -def apply_lora_mlp_geglu(self, X): +from .geglu import geglu_exact_forward_kernel, geglu_exact_backward_kernel +def apply_lora_mlp_geglu_exact(self, X): gateW, gateW_quant, gateA, gateB, gateS = get_lora_parameters(self.gate_proj) upW, upW_quant, upA, upB, upS = get_lora_parameters(self. up_proj) downW, downW_quant, downA, downB, downS = get_lora_parameters(self.down_proj) @@ -192,7 +192,21 @@ def apply_lora_mlp_geglu(self, X): gateW, gateW_quant, gateA, gateB, gateS, upW, upW_quant, upA, upB, upS, downW, downW_quant, downA, downB, downS, - geglu_forward_kernel, geglu_backward_kernel,) + geglu_exact_forward_kernel, geglu_exact_backward_kernel,) + return out +pass + + +from .geglu import geglu_approx_forward_kernel, geglu_approx_backward_kernel +def apply_lora_mlp_geglu_approx(self, X): + gateW, gateW_quant, gateA, gateB, gateS = get_lora_parameters(self.gate_proj) + upW, upW_quant, upA, upB, upS = get_lora_parameters(self. up_proj) + downW, downW_quant, downA, downB, downS = get_lora_parameters(self.down_proj) + out = LoRA_MLP.apply(X, + gateW, gateW_quant, gateA, gateB, gateS, + upW, upW_quant, upA, upB, upS, + downW, downW_quant, downA, downB, downS, + geglu_approx_forward_kernel, geglu_approx_backward_kernel,) return out pass diff --git a/unsloth/kernels/geglu.py b/unsloth/kernels/geglu.py index 0a952a10a9..1d29db8f5b 100644 --- a/unsloth/kernels/geglu.py +++ b/unsloth/kernels/geglu.py @@ -38,7 +38,7 @@ def _exact_forward_kernel(e, g, h, n_elements, BLOCK_SIZE : tl.constexpr,): pass -def geglu_forward_kernel(gate, up): +def geglu_exact_forward_kernel(gate, up): batch, seq_len, hd = gate.shape n_elements = gate.numel() out = torch.empty((batch, seq_len, hd), dtype = gate.dtype, device = "cuda") @@ -95,7 +95,7 @@ def _exact_backward_kernel(DW, e, g, n_elements, BLOCK_SIZE : tl.constexpr,): pass -def geglu_backward_kernel(DW, e, g): +def geglu_exact_backward_kernel(DW, e, g): batch_seq_len, hd = e.shape n_elements = e.numel() grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),) @@ -171,7 +171,7 @@ def _approx_backward_kernel(DW, e, g, n_elements, BLOCK_SIZE : tl.constexpr,): T = 1.0 + tl.math.tanh(a + b) T2 = 0.5 * T # Q = -T * (T - 2.0) * (a + 3.0 * b) - Q2 = -T2 * (T - 2.0) * (a + 3.0 * b) + Q2 = -T2 * (T - 2.0) * (a + 3.0 * b) df_de = T2 + Q2 # 1/2 * (T + Q) # f = 1/2 * e * (1 + tanh( sqrt(2/pi) * (x + 0.044715 * x^3 ) )) @@ -192,3 +192,12 @@ def _approx_backward_kernel(DW, e, g, n_elements, BLOCK_SIZE : tl.constexpr,): tl.store(e + offsets, df_row, mask = mask) # df = DW * f tl.store(g + offsets, de_row, mask = mask) # de pass + + +def geglu_approx_backward_kernel(DW, e, g): + batch_seq_len, hd = e.shape + n_elements = e.numel() + grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),) + _approx_backward_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,) + return DW, e, g +pass diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 54e016d1af..87ca5d8487 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -1335,7 +1335,7 @@ class FastLlamaModel: if model_type == "llama": apply_lora_mlp = apply_lora_mlp_swiglu elif model_type == "mistral": apply_lora_mlp = apply_lora_mlp_swiglu - elif model_type == "gemma": apply_lora_mlp = apply_lora_mlp_geglu + elif model_type == "gemma": apply_lora_mlp = apply_lora_mlp_geglu_approx else: raise NotImplementedError(f"Unsloth: {model_type} is not yet implemented!") pass