Update gemma.py
This commit is contained in:
parent
b0b38f770b
commit
40169f9fd3
1 changed files with 12 additions and 0 deletions
|
|
@ -97,6 +97,18 @@ class FastGemmaRotaryEmbedding(torch.nn.Module):
|
|||
self.cos_cached = emb.cos().to(dtype=x.dtype)
|
||||
self.sin_cached = emb.sin().to(dtype=x.dtype)
|
||||
pass
|
||||
|
||||
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)
|
||||
cos_cached = emb.cos().to(dtype=x.dtype)
|
||||
sin_cached = emb.sin().to(dtype=x.dtype)
|
||||
|
||||
print(cos_cached)
|
||||
print(self.cos_cached)
|
||||
|
||||
# return self.cos_cached[:,:length], self.sin_cached[:,:length]
|
||||
return self.cos_cached, self.sin_cached
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue