inference

This commit is contained in:
Daniel Han-Chen 2024-02-01 17:20:02 +11:00
commit cf4b58eeb6
2 changed files with 5 additions and 7 deletions

View file

@ -308,8 +308,6 @@ def LlamaAttention_fast_forward(
padding_mask: Optional[torch.LongTensor] = None,
*args, **kwargs,
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
bsz, q_len, _ = hidden_states.size()
# Check for inference
if past_key_value is not None:
@ -322,6 +320,8 @@ def LlamaAttention_fast_forward(
return A, None, past_key_value
pass
bsz, q_len, _ = hidden_states.size()
n_heads = self.num_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
@ -430,7 +430,6 @@ def LlamaDecoderLayer_fast_forward(
(see `past_key_values`).
past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
"""
bsz, q_len, hd = hidden_states.size()
if past_key_value is not None:
# Self Attention
residual = hidden_states
@ -658,8 +657,7 @@ def LlamaModel_fast_forward(
if output_attentions:
all_self_attns += (layer_outputs[1],)
pass
bsz, q_len, hd = hidden_states.size()
if past_key_values is not None:
hidden_states = fast_rms_layernorm_inference(self.norm, hidden_states)
else:

View file

@ -46,8 +46,6 @@ def MistralAttention_fast_forward(
*args, **kwargs,
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
bsz, q_len, _ = hidden_states.size()
# Check for inference
if past_key_value is not None:
A, past_key_value = LlamaAttention_fast_forward_inference(
@ -59,6 +57,8 @@ def MistralAttention_fast_forward(
return A, None, past_key_value
pass
bsz, q_len, _ = hidden_states.size()
n_heads = self.num_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads