diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index 154049e473..2e5ddaa9ca 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -73,19 +73,14 @@ class FastGemmaRotaryEmbedding(torch.nn.Module): self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device="cuda").float() / self.dim) ) - t = torch.arange(position_ids.shape[1], device="cuda", dtype=torch.int64).float().unsqueeze(0) + t = torch.arange(self.max_position_embeddings, device="cpu", dtype=torch.int64).float().to("cuda").unsqueeze(0) inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to("cuda") position_ids_expanded = t[:, None, :].float() freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2) - emb_old = torch.cat((freqs, freqs), dim=-1) + emb = torch.cat((freqs, freqs), dim=-1) - inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to("cuda") - position_ids_expanded = position_ids[:, None, :].float() - freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2) - emb_new = torch.cat((freqs, freqs), dim=-1) - - logger.warning_once(str(torch.dist(emb_old, emb_new))) - return emb_old.cos().to(dtype=x.dtype), emb_old.sin().to(dtype=x.dtype) + seq_len = position_ids.shape[1] + return emb.cos().to(dtype=x.dtype)[:seq_len], emb.sin().to(dtype=x.dtype)[:seq_len] pass