Update gemma2.py

This commit is contained in:
Daniel Han 2024-11-16 15:18:38 -08:00
commit e7ad484169

View file

@ -60,8 +60,7 @@ if HAS_FLASH_ATTENTION_SOFTCAPPING:
from flash_attn import flash_attn_func
# [TODO] We must randomnly use torch.compile?
# I checked the gradients and formulas and I'm sure it's correct.
# I'm stumped :(
# Gemma 2 uses double RMS Layernorms, so the backward passes should not overwrite the gradients!
@torch.compile(fullgraph = False, dynamic = True, options = torch_compile_options)
def fast_rms_layernorm_gemma2_compiled(layernorm, X, gemma = True):
old_dtype = X.dtype