diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 75d491cb1a..43924d0ab3 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -214,3 +214,49 @@ def unpatch_rms_layernorm(): transformers.models.llama.modeling_llama.LlamaRMSNorm = LlamaRMSNorm return pass + + +def test_rms_layernorm( + dim = 1024, eps = 1e-5, dtype = torch.float16, + bsz = 21, random_state = 3407, seqlen = 3341, +): + from transformers.models.llama.modeling_llama import LlamaRMSNorm + layernorm = LlamaRMSNorm((dim,), eps = eps).to("cuda") + torch.cuda.manual_seed(random_state) + torch.manual_seed(random_state) + torch.nn.init.uniform_(layernorm.weight) + X = torch.randn((bsz, seqlen, dim), dtype = dtype, device = "cuda") + XX = X.clone() + X .requires_grad_(True) + XX.requires_grad_(True) + Y = layernorm(X) + YY = torch.randn((bsz, seqlen, dim), dtype = dtype, device = "cuda", requires_grad = True) + Y.backward(YY) + correct_grad = X.grad.clone() + # from unsloth.kernels import fast_rms_layernorm + Y = fast_rms_layernorm(layernorm, XX) + Y.backward(YY) + assert(torch.amax(correct_grad - XX.grad).item() <= 0.05) +pass + + +def testing_suite_layernorm(): + for dim in [512, 1024, 2048]: + for dtype in [torch.float16, torch.bfloat16]: + with torch.autocast(device_type = "cuda", dtype = dtype): + for seqlen in [3341, 2048, 349]: + for random_state in [3407, 42]: + test_rms_layernorm( + dim = dim, + eps = 1e-5, + dtype = dtype, + bsz = 21, + random_state = random_state, + seqlen = seqlen, + ) + pass + pass + pass + pass + pass +pass