Approx gelu
This commit is contained in:
parent
1c7f0d21ee
commit
786885c38c
3 changed files with 30 additions and 7 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue