From 74382dea479aecb830d553da8c2edef1f66f5b3b Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sat, 16 Nov 2024 12:18:47 -0800 Subject: [PATCH] Update rms_layernorm.py --- unsloth/kernels/rms_layernorm.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) 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),