diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 3b54604a0b..4b22f8c3e5 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -57,7 +57,6 @@ 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, @@ -79,9 +78,6 @@ 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) @@ -95,7 +91,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(dX + col_offsets, output, mask = mask) + tl.store(dY + col_offsets, output, mask = mask) pass @@ -176,11 +172,9 @@ 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),