From da6de6dcb84f99180ed191182eaedbf5c9c4f968 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 24 Sep 2024 22:50:33 -0700 Subject: [PATCH] RMS Layernorm --- unsloth/kernels/__init__.py | 6 +++++- unsloth/kernels/rms_layernorm.py | 22 ++++++++++++++++++++++ 2 files changed, 27 insertions(+), 1 deletion(-) diff --git a/unsloth/kernels/__init__.py b/unsloth/kernels/__init__.py index 841d0ce0f0..606adf80f7 100644 --- a/unsloth/kernels/__init__.py +++ b/unsloth/kernels/__init__.py @@ -13,7 +13,11 @@ # limitations under the License. from .cross_entropy_loss import fast_cross_entropy_loss -from .rms_layernorm import fast_rms_layernorm +from .rms_layernorm import ( + fast_rms_layernorm, + patch_rms_layernorm, + unpatch_rms_layernorm, +) from .layernorm import ( fast_layernorm, patch_layernorm, diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index ac5beb5ab1..75d491cb1a 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -192,3 +192,25 @@ def fast_rms_layernorm(layernorm, X, gemma = False): out = Fast_RMS_Layernorm.apply(X, W, eps, gemma) return out pass + + +from transformers.models.llama.modeling_llama import LlamaRMSNorm +class Unsloth_LlamaRMSNorm(LlamaRMSNorm): + def forward(self, X): + return fast_rms_layernorm(self, X, gemma = False) + pass +pass + + +def patch_rms_layernorm(): + import transformers.models.llama.modeling_llama + transformers.models.llama.modeling_llama.LlamaRMSNorm = Unsloth_LlamaRMSNorm + return +pass + + +def unpatch_rms_layernorm(): + import transformers.models.llama.modeling_llama + transformers.models.llama.modeling_llama.LlamaRMSNorm = LlamaRMSNorm + return +pass