fast inference
This commit is contained in:
parent
257cd7d531
commit
31578f2010
2 changed files with 29 additions and 99 deletions
|
|
@ -78,10 +78,8 @@ def LlamaAttention_fast_forward_inference(
|
|||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
past_key_value: Optional[Tuple[torch.Tensor]],
|
||||
position_ids,
|
||||
do_prefill = False,
|
||||
temp_QA = None,
|
||||
temp_KV = None,
|
||||
RH_Q = None,
|
||||
):
|
||||
"""
|
||||
https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L406
|
||||
|
|
@ -132,6 +130,9 @@ def LlamaAttention_fast_forward_inference(
|
|||
self.paged_attention_V = self.paged_attention[:,1]
|
||||
self.paged_attention_K[:seq_len] = K1.permute(2, 0, 1, 3)
|
||||
self.paged_attention_V[:seq_len] = V1.permute(2, 0, 1, 3)
|
||||
self.temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda")
|
||||
self.temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda")
|
||||
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda")
|
||||
self.attention = torch.empty((bsz, n_heads, 1, KV_CACHE_INCREMENT), dtype = dtype, device = "cuda")
|
||||
self.scalar = 1.0 / math_sqrt(self.head_dim)
|
||||
elif kv_seq_len >= self.paged_attention.shape[0]:
|
||||
|
|
@ -141,15 +142,9 @@ def LlamaAttention_fast_forward_inference(
|
|||
self.attention.resize_((bsz, n_heads, 1, self.attention.shape[-1]+KV_CACHE_INCREMENT))
|
||||
pass
|
||||
|
||||
if temp_QA is None:
|
||||
temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda")
|
||||
temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda")
|
||||
RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda")
|
||||
pass
|
||||
|
||||
Qn = fast_linear_forward(self.q_proj, Xn, out = temp_QA[0])
|
||||
Kn = fast_linear_forward(self.k_proj, Xn, out = temp_KV[0])
|
||||
Vn = fast_linear_forward(self.v_proj, Xn, out = temp_KV[1])
|
||||
Qn = fast_linear_forward(self.q_proj, Xn, out = self.temp_QA[0])
|
||||
Kn = fast_linear_forward(self.k_proj, Xn, out = self.temp_KV[0])
|
||||
Vn = fast_linear_forward(self.v_proj, Xn, out = self.temp_KV[1])
|
||||
Qn = Qn.view(bsz, 1, n_heads, head_dim).transpose(1, 2)
|
||||
Kn = Kn.view(bsz, 1, n_kv_heads, head_dim).transpose(1, 2)
|
||||
Vn = Vn.view(bsz, 1, n_kv_heads, head_dim).transpose(1, 2)
|
||||
|
|
@ -160,7 +155,7 @@ def LlamaAttention_fast_forward_inference(
|
|||
sin = self.rotary_emb.sin_cached[seq_len]
|
||||
h = head_dim // 2
|
||||
|
||||
# RH_Q = self.RH_Q
|
||||
RH_Q = self.RH_Q
|
||||
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);
|
||||
|
||||
|
|
@ -193,19 +188,17 @@ def LlamaAttention_fast_forward_inference(
|
|||
A = torch.matmul(A, Vnn, out = Qn)
|
||||
A = A.transpose(1, 2)
|
||||
A = A.reshape(bsz, 1, self.hidden_size)
|
||||
A = fast_linear_forward(self.o_proj, A, out = temp_QA[1])
|
||||
A = fast_linear_forward(self.o_proj, A, out = self.temp_QA[1])
|
||||
return A, (Kn, Vn)
|
||||
pass
|
||||
|
||||
|
||||
def fast_mlp_inference(self, X, temp = None):
|
||||
def fast_mlp_inference(self, X):
|
||||
# gate = self.gate_proj(X)
|
||||
# up = self.up_proj(X)
|
||||
bsz, _, hd = X.shape
|
||||
if temp is None:
|
||||
mlp_size = self.config.intermediate_size
|
||||
temp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda")
|
||||
pass
|
||||
mlp_size = self.config.intermediate_size
|
||||
temp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda")
|
||||
|
||||
gate = fast_linear_forward(self.gate_proj, X, out = temp[0])
|
||||
up = fast_linear_forward(self. up_proj, X, out = temp[1])
|
||||
|
|
@ -218,26 +211,13 @@ def fast_mlp_inference(self, X, temp = None):
|
|||
pass
|
||||
|
||||
|
||||
def fast_rms_layernorm_inference(self, X, temp1 = None, temp2 = None):
|
||||
def fast_rms_layernorm_inference(self, X):
|
||||
old_dtype = X.dtype
|
||||
|
||||
if temp1 is None:
|
||||
bsz, _, hd = X.shape
|
||||
temp1 = torch.empty((2, bsz, 1, hd), dtype = torch.float32, device = "cuda")
|
||||
temp2 = torch.empty((bsz, 1, hd), dtype = old_dtype, device = "cuda")
|
||||
pass
|
||||
|
||||
XX = temp1[0]
|
||||
XX2 = temp1[1]
|
||||
|
||||
# XX = X.to(torch.float32)
|
||||
XX[:] = X
|
||||
# variance = XX.square().mean(-1, keepdim = True)
|
||||
variance = torch.square(XX, out = XX2).mean(-1, keepdim = True)
|
||||
XX = X.to(torch.float32)
|
||||
variance = XX.square().mean(-1, keepdim = True)
|
||||
variance += self.variance_epsilon
|
||||
XX *= variance.rsqrt_()
|
||||
# X = XX.to(old_dtype) # Must preserve due to residual
|
||||
temp2[:] = XX; X = temp2
|
||||
X = XX.to(old_dtype) # Must preserve due to residual
|
||||
X *= self.weight
|
||||
return X
|
||||
pass
|
||||
|
|
@ -262,6 +242,9 @@ def LlamaAttention_fast_forward(
|
|||
del self.paged_attention_K
|
||||
del self.paged_attention_V
|
||||
del self.paged_attention
|
||||
del self.temp_QA
|
||||
del self.temp_KV
|
||||
del self.RH_Q
|
||||
del self.attention
|
||||
pass
|
||||
|
||||
|
|
@ -385,6 +368,7 @@ def LlamaDecoderLayer_fast_forward(
|
|||
self.self_attn,
|
||||
hidden_states,
|
||||
past_key_value,
|
||||
position_ids,
|
||||
do_prefill = do_prefill,
|
||||
)
|
||||
hidden_states += residual
|
||||
|
|
@ -443,16 +427,6 @@ def LlamaModel_fast_forward(
|
|||
return_dict: Optional[bool] = None,
|
||||
*args, **kwargs,
|
||||
) -> Union[Tuple, BaseModelOutputWithPast]:
|
||||
|
||||
# Clear inference
|
||||
if hasattr(self, "temp_QA"):
|
||||
del self.temp_QA
|
||||
del self.temp_KV
|
||||
del self.RH_Q
|
||||
del self.mlp_temp
|
||||
del self.layernorm1
|
||||
del self.layernorm2
|
||||
pass
|
||||
|
||||
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||
assert(output_attentions is False)
|
||||
|
|
@ -638,78 +612,31 @@ def LlamaModel_fast_forward_inference(
|
|||
):
|
||||
# Fix out of bounds tokenization
|
||||
input_ids = input_ids[:,:self.max_seq_length]
|
||||
|
||||
hidden_states = self.embed_tokens(input_ids)
|
||||
bsz, q_len, hd = hidden_states.shape
|
||||
dtype = hidden_states.dtype
|
||||
|
||||
first_attention = self.layers[0].self_attn
|
||||
n_heads = first_attention.num_heads
|
||||
n_groups = first_attention.num_key_value_groups
|
||||
n_kv_heads = first_attention.num_key_value_heads
|
||||
head_dim = first_attention.head_dim
|
||||
mlp_size = self.config.intermediate_size
|
||||
|
||||
# Temporary matrices
|
||||
if not hasattr(self, "temp_QA"):
|
||||
self.temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda")
|
||||
self.temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda")
|
||||
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda")
|
||||
self.mlp_temp = torch.empty((2, bsz, 1, mlp_size), dtype = dtype, device = "cuda")
|
||||
self.layernorm1 = torch.empty((2, bsz, 1, hd), dtype = torch.float32, device = "cuda")
|
||||
self.layernorm2 = torch.empty((bsz, 1, hd), dtype = dtype, device = "cuda")
|
||||
pass
|
||||
temp_QA = self.temp_QA
|
||||
temp_KV = self.temp_KV
|
||||
RH_Q = self.RH_Q
|
||||
mlp_temp = self.mlp_temp
|
||||
layernorm1 = self.layernorm1
|
||||
layernorm2 = self.layernorm2
|
||||
#
|
||||
|
||||
next_decoder_cache = []
|
||||
for idx, decoder_layer in enumerate(self.layers):
|
||||
# Self Attention
|
||||
residual = hidden_states
|
||||
hidden_states = fast_rms_layernorm_inference(
|
||||
decoder_layer.input_layernorm,
|
||||
hidden_states,
|
||||
layernorm1,
|
||||
layernorm2,
|
||||
)
|
||||
hidden_states = fast_rms_layernorm_inference(decoder_layer.input_layernorm, hidden_states)
|
||||
hidden_states, present_key_value = LlamaAttention_fast_forward_inference(
|
||||
decoder_layer.self_attn,
|
||||
hidden_states,
|
||||
past_key_values[idx],
|
||||
do_prefill = False,
|
||||
temp_QA = temp_QA,
|
||||
temp_KV = temp_KV,
|
||||
RH_Q = RH_Q,
|
||||
None,
|
||||
)
|
||||
hidden_states += residual
|
||||
|
||||
# Fully Connected
|
||||
residual = hidden_states
|
||||
hidden_states = fast_rms_layernorm_inference(
|
||||
decoder_layer.post_attention_layernorm,
|
||||
hidden_states,
|
||||
layernorm1,
|
||||
layernorm2,
|
||||
)
|
||||
hidden_states = fast_mlp_inference(
|
||||
decoder_layer.mlp,
|
||||
hidden_states,
|
||||
mlp_temp,
|
||||
)
|
||||
hidden_states = fast_rms_layernorm_inference(decoder_layer.post_attention_layernorm, hidden_states)
|
||||
hidden_states = fast_mlp_inference(decoder_layer.mlp, hidden_states)
|
||||
hidden_states += residual
|
||||
|
||||
next_decoder_cache.append(present_key_value)
|
||||
pass
|
||||
hidden_states = fast_rms_layernorm_inference(
|
||||
self.norm,
|
||||
hidden_states,
|
||||
layernorm1,
|
||||
layernorm2,
|
||||
)
|
||||
hidden_states = fast_rms_layernorm_inference(self.norm, hidden_states)
|
||||
|
||||
return BaseModelOutputWithPast(
|
||||
last_hidden_state = hidden_states,
|
||||
|
|
|
|||
|
|
@ -51,6 +51,9 @@ def MistralAttention_fast_forward(
|
|||
del self.paged_attention_K
|
||||
del self.paged_attention_V
|
||||
del self.paged_attention
|
||||
del self.temp_QA
|
||||
del self.temp_KV
|
||||
del self.RH_Q
|
||||
del self.attention
|
||||
pass
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue