Update gemma.py

This commit is contained in:
Daniel Han-Chen 2024-02-25 04:22:35 +11:00
commit 6ad44835d2

View file

@ -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