From 5adc84720291b8b8d2882d0002d148e86e29643f Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sat, 16 Nov 2024 13:55:11 -0800 Subject: [PATCH] Gemma --- unsloth/kernels/rms_layernorm.py | 8 +++++++- unsloth/models/gemma2.py | 8 ++++---- 2 files changed, 11 insertions(+), 5 deletions(-) diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 4b22f8c3e5..3b54604a0b 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -57,6 +57,7 @@ pass @triton.jit def _rms_layernorm_backward( dY, dY_row_stride, + dX, dX_row_stride, X, X_row_stride, W, W_row_stride, r, r_row_stride, @@ -78,6 +79,9 @@ def _rms_layernorm_backward( X += row_idx * X_row_stride r += row_idx * r_row_stride + if GEMMA: dX += row_idx * dY_row_stride + else: dX = dY + dY_row = tl.load(dY + col_offsets, mask = mask, other = 0).to(tl.float32) X_row = tl.load(X + col_offsets, mask = mask, other = 0).to(tl.float32) W_row = tl.load(W + col_offsets, mask = mask, other = 0).to(tl.float32) @@ -91,7 +95,7 @@ def _rms_layernorm_backward( rowsum_dY_normed = tl.sum(dY_W * normed, axis = 0) output = inv_var/n_cols * (n_cols*dY_W - normed*rowsum_dY_normed) - tl.store(dY + col_offsets, output, mask = mask) + tl.store(dX + col_offsets, output, mask = mask) pass @@ -172,9 +176,11 @@ class Fast_RMS_Layernorm(torch.autograd.Function): n_cols : int n_rows, n_cols = dY.shape # dW = X + dX = torch.empty_like(dY, device = "cuda:0") if ctx.GEMMA else dY _rms_layernorm_backward[(n_rows,)]( dY, dY.stride(0), + dX, dX.stride(0), X, X .stride(0), W, W .stride(0), r, r .stride(0), diff --git a/unsloth/models/gemma2.py b/unsloth/models/gemma2.py index 4eb9d64313..872824b11a 100644 --- a/unsloth/models/gemma2.py +++ b/unsloth/models/gemma2.py @@ -207,7 +207,7 @@ def Gemma2DecoderLayer_fast_forward( hidden_states += residual else: residual = hidden_states - hidden_states = fast_rms_layernorm_gemma2_compiled(self.input_layernorm, hidden_states, gemma = True) + hidden_states = fast_rms_layernorm(self.input_layernorm, hidden_states, gemma = True) hidden_states, self_attn_weights, present_key_value = self.self_attn( hidden_states=hidden_states, causal_mask=causal_mask, @@ -218,14 +218,14 @@ def Gemma2DecoderLayer_fast_forward( use_cache=use_cache, padding_mask=padding_mask, ) - hidden_states = fast_rms_layernorm_gemma2_compiled(self.post_attention_layernorm, hidden_states, gemma = True) + hidden_states = fast_rms_layernorm(self.post_attention_layernorm, hidden_states, gemma = True) hidden_states = residual + hidden_states # Fully Connected residual = hidden_states - hidden_states = fast_rms_layernorm_gemma2_compiled(self. pre_feedforward_layernorm, hidden_states, gemma = True) + hidden_states = fast_rms_layernorm(self. pre_feedforward_layernorm, hidden_states, gemma = True) hidden_states = self.mlp(hidden_states) - hidden_states = fast_rms_layernorm_gemma2_compiled(self.post_feedforward_layernorm, hidden_states, gemma = True) + hidden_states = fast_rms_layernorm(self.post_feedforward_layernorm, hidden_states, gemma = True) hidden_states = residual + hidden_states pass