From dfb2e250f6f332a3d8aa49da2391d928c535bdd9 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 25 Feb 2024 02:17:12 +1100 Subject: [PATCH] Update gemma.py --- unsloth/models/gemma.py | 28 +++++++++++++++++++++++++++- 1 file changed, 27 insertions(+), 1 deletion(-) diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index 301f7fa923..22734a0933 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -20,6 +20,7 @@ from transformers.models.gemma.modeling_gemma import ( GemmaDecoderLayer, GemmaModel, GemmaForCausalLM, + GemmaRotaryEmbedding, apply_rotary_pos_emb, repeat_kv, ) @@ -56,6 +57,31 @@ def fast_geglu_inference(self, X): pass +class FastGemmaRotaryEmbedding(nn.Module): + def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None): + super().__init__() + + self.dim = dim + self.max_position_embeddings = max_position_embeddings + self.base = base + self.register_buffer("inv_freq", None, persistent=False) + + def forward(self, x, position_ids, seq_len=None): + # x: [bs, num_attention_heads, seq_len, head_size] + if self.inv_freq is None: + self.inv_freq = 1.0 / ( + self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device=x.device).float() / self.dim) + ) + print(position_ids) + + inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) + position_ids_expanded = position_ids[:, None, :].float() + freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2) + emb = torch.cat((freqs, freqs), dim=-1) + return emb.cos().to(dtype=x.dtype), emb.sin().to(dtype=x.dtype) +pass + + # https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L320 def GemmaAttention_fast_forward( self, @@ -513,7 +539,7 @@ class FastGemmaModel(FastLlamaModel): GemmaModel .forward = GemmaModel_fast_forward GemmaForCausalLM .forward = GemmaForCausalLM_fast_forward PeftModelForCausalLM.forward = PeftModelForCausalLM_fast_forward - + GemmaRotaryEmbedding = FastGemmaRotaryEmbedding # Solves https://github.com/unslothai/unsloth/issues/168 # Static KV Cache was introduced in 4.38.0, causing training to be much slower. # Inferene can now be CUDAGraphed, but we shall retain the old rotary embeddings.