From 5bf8cacebd67fa40b7a4d150414722d262bb0fb1 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 25 Feb 2024 01:58:33 +1100 Subject: [PATCH] Update gemma.py --- unsloth/models/gemma.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index 8f173e210f..646aec488d 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -97,7 +97,7 @@ def GemmaAttention_fast_forward( sin = self.rotary_emb.sin_cached Q, K = fast_rope_embedding(Q, K, cos, sin) else: - cos, sin = self.rotary_emb(V, position_ids, seq_len=None) + cos, sin = self.rotary_emb(V, position_ids, seq_len = q_len+1) Q, K = inplace_rope_embedding(Q, K, cos, sin, position_ids) pass @@ -518,8 +518,8 @@ class FastGemmaModel(FastLlamaModel): # Inferene can now be CUDAGraphed, but we shall retain the old rotary embeddings. # https://github.com/huggingface/transformers/pull/27931 # https://github.com/huggingface/transformers/blob/v4.37.2/src/transformers/models/llama/modeling_llama.py - # import transformers.models.gemma.modeling_gemma - # transformers.models.gemma.modeling_gemma.GemmaRotaryEmbedding = LlamaRotaryEmbedding + import transformers.models.gemma.modeling_gemma + transformers.models.gemma.modeling_gemma.GemmaRotaryEmbedding = LlamaRotaryEmbedding return pass