Update gemma.py

This commit is contained in:
Daniel Han-Chen 2024-02-25 03:03:23 +11:00
commit 10d03d8315

View file

@ -73,10 +73,7 @@ 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).unsqueeze(0)
print(t)
print(position_ids)
raise False
t = torch.arange(position_ids.shape[1], device="cuda", dtype=torch.int64).float().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)
@ -87,6 +84,7 @@ class FastGemmaRotaryEmbedding(torch.nn.Module):
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_new.cos().to(dtype=x.dtype), emb_new.sin().to(dtype=x.dtype)
pass