Update llama.py

This commit is contained in:
Daniel Han-Chen 2024-02-01 23:38:07 +11:00
commit 73f63d6884

View file

@ -250,7 +250,7 @@ def LlamaAttention_fast_forward_inference(
sin = self.rotary_emb.sin_cached[seq_len]
h = head_dim // 2
RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = Xn.dtype, device = "cuda")
RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda")
RH_Q[:,:,:,:h] = Qn[:,:,:,h:]; RH_Q[:,:,:,h:] = Qn[:,:,:,:h]; torch.neg(RH_Q[:,:,:,:h], out = RH_Q[:,:,:,:h]);
Qn *= cos; Qn.addcmul_(RH_Q, sin);
@ -259,8 +259,17 @@ def LlamaAttention_fast_forward_inference(
Kn *= cos; Kn.addcmul_(RH_K, sin);
# New KV cache
Kn = torch.cat([K1, Kn], dim = 2)
Vn = torch.cat([V1, Vn], dim = 2)
# Kn = torch.cat([K1, Kn], dim = 2)
# Vn = torch.cat([V1, Vn], dim = 2)
if not hasattr(self, "paged_attention_K"):
paged_attention = torch.empty((2, bsz, n_kv_heads, 2048, head_dim), dtype = dtype, device = "cuda")
self.paged_attention_K = paged_attention[0]
self.paged_attention_V = paged_attention[1]
self.paged_attention_K[:,:,:kv_seq_len,:] = K1
self.paged_attention_V[:,:,:kv_seq_len,:] = V1
pass
Kn = self.paged_attention_K[:,:,seq_len,:] = K1
Vn = self.paged_attention_V[:,:,seq_len,:] = V1
# Grouped query attention
if n_groups != 1: