From 6ad44835d29879f019580a4383442338fa620367 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 25 Feb 2024 04:22:35 +1100 Subject: [PATCH] Update gemma.py --- unsloth/models/gemma.py | 18 ++++++++++++------ 1 file changed, 12 insertions(+), 6 deletions(-) diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index 0045e6b712..ebb242b5db 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -65,6 +65,8 @@ class FastGemmaRotaryEmbedding(torch.nn.Module): self.max_position_embeddings = max_position_embeddings self.base = base self.register_buffer("inv_freq", None, persistent=False) + self.register_buffer("cos_cached", None, persistent=False) + self.register_buffer("sin_cached", None, persistent=False) def forward(self, x, position_ids, seq_len=None): # x: [bs, num_attention_heads, seq_len, head_size] @@ -73,13 +75,17 @@ class FastGemmaRotaryEmbedding(torch.nn.Module): self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device=x.device).float() / self.dim) ) - position_ids = torch.arange(self.max_position_embeddings, device=x.device, dtype=torch.int64).unsqueeze(0) - inv_freq_expanded = self.inv_freq[None, :, None].float().expand(1, -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) length = position_ids.shape[1] - return emb.cos().to(dtype=x.dtype)[:,:length], emb.sin().to(dtype=x.dtype)[:,:length] + if self.cos_cached is None: + position_ids = torch.arange(self.max_position_embeddings, device=x.device, dtype=torch.int64).unsqueeze(0) + inv_freq_expanded = self.inv_freq[None, :, None].float().expand(1, -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) + self.cos_cached = emb.cos().to(dtype=x.dtype) + self.sin_cached = emb.sin().to(dtype=x.dtype) + pass + return self.cos_cached[:,:length], self.sin_cached[:,:length] pass