Update rms_layernorm.py
This commit is contained in:
parent
fc4ca43ee1
commit
634489002d
1 changed files with 46 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue