correct_dtype

This commit is contained in:
Daniel Han-Chen 2024-02-26 14:39:43 +11:00
commit 7a5db6758c
2 changed files with 6 additions and 6 deletions

View file

@ -630,10 +630,10 @@ class FastGemmaModel(FastLlamaModel):
pass
# Downcast RoPE embedding to correct data type
if (name.endswith("rotary_emb") or hasattr(module, "cos_cached")) \
and (module.cos_cached.dtype != expected_dtype):
and (module.cos_cached.dtype != correct_dtype):
module.cos_cached = module.cos_cached.to(expected_dtype)
module.sin_cached = module.sin_cached.to(expected_dtype)
module.cos_cached = module.cos_cached.to(correct_dtype)
module.sin_cached = module.sin_cached.to(correct_dtype)
pass
pass
pass

View file

@ -1172,10 +1172,10 @@ class FastLlamaModel:
pass
# Downcast RoPE embedding to correct data type
if (name.endswith("rotary_emb") or hasattr(module, "cos_cached")) \
and (module.cos_cached.dtype != expected_dtype):
and (module.cos_cached.dtype != correct_dtype):
module.cos_cached = module.cos_cached.to(expected_dtype)
module.sin_cached = module.sin_cached.to(expected_dtype)
module.cos_cached = module.cos_cached.to(correct_dtype)
module.sin_cached = module.sin_cached.to(correct_dtype)
pass
pass
pass