From 045bce2fa56ed8ad4d495553f61bc4708f00d37b Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 25 Feb 2024 02:25:39 +1100 Subject: [PATCH] Update gemma.py --- unsloth/models/gemma.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index 97528ebbce..aa87d846b0 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -72,12 +72,15 @@ class FastGemmaRotaryEmbedding(torch.nn.Module): 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) + t = torch.arange(self.max_position_embeddings, device="cpu", dtype=torch.int64).float() + freqs = torch.outer(t, self.inv_freq) emb = torch.cat((freqs, freqs), dim=-1) + + # 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