Approx gelu
This commit is contained in:
parent
e032445694
commit
c88ab10a5c
2 changed files with 41 additions and 5 deletions
|
|
@ -19,7 +19,7 @@ from .utils import calculate_settings
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def _forward_kernel(e, g, h, n_elements, BLOCK_SIZE : tl.constexpr,):
|
def _exact_forward_kernel(e, g, h, n_elements, BLOCK_SIZE : tl.constexpr,):
|
||||||
block_idx = tl.program_id(0)
|
block_idx = tl.program_id(0)
|
||||||
offsets = block_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
offsets = block_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
mask = offsets < n_elements
|
mask = offsets < n_elements
|
||||||
|
|
@ -43,13 +43,13 @@ def geglu_forward_kernel(gate, up):
|
||||||
n_elements = gate.numel()
|
n_elements = gate.numel()
|
||||||
out = torch.empty((batch, seq_len, hd), dtype = gate.dtype, device = "cuda")
|
out = torch.empty((batch, seq_len, hd), dtype = gate.dtype, device = "cuda")
|
||||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||||
_forward_kernel[grid](gate, up, out, n_elements, BLOCK_SIZE = 1024,)
|
_exact_forward_kernel[grid](gate, up, out, n_elements, BLOCK_SIZE = 1024,)
|
||||||
return out
|
return out
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def _backward_kernel(DW, e, g, n_elements, BLOCK_SIZE : tl.constexpr,):
|
def _exact_backward_kernel(DW, e, g, n_elements, BLOCK_SIZE : tl.constexpr,):
|
||||||
"""
|
"""
|
||||||
f = 1/2 * e * (1 + erf(1/sqrt(2) * e))
|
f = 1/2 * e * (1 + erf(1/sqrt(2) * e))
|
||||||
h = f * up
|
h = f * up
|
||||||
|
|
@ -99,6 +99,42 @@ def geglu_backward_kernel(DW, e, g):
|
||||||
batch_seq_len, hd = e.shape
|
batch_seq_len, hd = e.shape
|
||||||
n_elements = e.numel()
|
n_elements = e.numel()
|
||||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||||
_backward_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
_exact_backward_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
||||||
return DW, e, g
|
return DW, e, g
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _approx_forward_kernel(e, g, h, n_elements, BLOCK_SIZE : tl.constexpr,):
|
||||||
|
block_idx = tl.program_id(0)
|
||||||
|
offsets = block_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offsets < n_elements
|
||||||
|
|
||||||
|
# f = 1/2 * e * (1 + tanh( sqrt(2/pi) * (x + 0.044715 * x^3 ) ))
|
||||||
|
# f = 1/2 * e * (1 + tanh( sqrt(2/pi) * x * (1 + 0.044715 * x^2 ) ))
|
||||||
|
# h = f * up
|
||||||
|
s = 0.7978845608028654 # math.sqrt(2 / math.pi)
|
||||||
|
|
||||||
|
e_row = tl.load(e + offsets, mask = mask, other = 0).to(tl.float32)
|
||||||
|
g_row = tl.load(g + offsets, mask = mask, other = 0)#.to(tl.float32)
|
||||||
|
|
||||||
|
f_row = 0.5 * e_row * (
|
||||||
|
tl.math.tanh(s * e_row * (1.0 + 0.044715 * e_row * e_row)) \
|
||||||
|
+ 1.0
|
||||||
|
)
|
||||||
|
f_row = f_row.to(g_row.dtype) # Exact copy from HF
|
||||||
|
h_row = f_row * g_row
|
||||||
|
|
||||||
|
# Store h
|
||||||
|
tl.store(h + offsets, h_row, mask = mask)
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def geglu_approx_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")
|
||||||
|
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||||
|
_approx_forward_kernel[grid](gate, up, out, n_elements, BLOCK_SIZE = 1024,)
|
||||||
|
return out
|
||||||
|
pass
|
||||||
|
|
|
||||||
|
|
@ -48,7 +48,7 @@ def fast_geglu_inference(self, X):
|
||||||
|
|
||||||
gate = fast_linear_forward(self.gate_proj, X, out = temp[0])
|
gate = fast_linear_forward(self.gate_proj, X, out = temp[0])
|
||||||
up = fast_linear_forward(self. up_proj, X, out = temp[1])
|
up = fast_linear_forward(self. up_proj, X, out = temp[1])
|
||||||
gate = torch.nn.functional.gelu(gate)
|
gate = torch.nn.functional.gelu(gate, approximate = "tanh")
|
||||||
gate *= up
|
gate *= up
|
||||||
|
|
||||||
# X = self.down_proj(gate)
|
# X = self.down_proj(gate)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue