Update llama.py

This commit is contained in:
Daniel Han-Chen 2024-02-02 20:19:15 +11:00
commit 82ead808d5

View file

@ -416,7 +416,6 @@ def LlamaDecoderLayer_fast_forward(
past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
"""
if past_key_value is not None and hasattr(self.self_attn, "paged_attention"):
print("1", end = "")
# Self Attention
residual = hidden_states
hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states)
@ -434,7 +433,6 @@ def LlamaDecoderLayer_fast_forward(
hidden_states = fast_mlp_inference(self.mlp, hidden_states)
hidden_states += residual
elif past_key_value is not None:
print("0", end = "")
# Self Attention
residual = hidden_states
hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states)
@ -680,6 +678,51 @@ def LlamaModel_fast_forward(
pass
# https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L825
@torch.compile
@torch.inference_mode
def LlamaModel_fast_forward_inference(
self,
input_ids
past_key_values,
):
# Fix out of bounds tokenization
input_ids = input_ids[:,:self.max_seq_length]
hidden_states = self.embed_tokens(input_ids)
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)
hidden_states, present_key_value = LlamaAttention_fast_forward_inference(
decoder_layer.self_attn,
hidden_states,
past_key_values[idx],
position_ids,
)
hidden_states += residual
# Fully Connected
residual = hidden_states
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)
return BaseModelOutputWithPast(
last_hidden_state = hidden_states,
past_key_values = next_decoder_cache,
hidden_states = [],
attentions = [],
)
pass
def LlamaForCausalLM_fast_forward(
self,
input_ids: torch.LongTensor = None,
@ -707,18 +750,29 @@ def LlamaForCausalLM_fast_forward(
# decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
self.model._has_no_labels = labels is None
outputs = self.model(
input_ids=input_ids,
causal_mask=causal_mask,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
if past_key_value is not None and \
hasattr(self.model.layers[0].self_attn, "paged_attention"):
outputs = LlamaModel_fast_forward_inference(
self.model,
input_ids,
past_key_values,
)
else:
outputs = self.model(
input_ids=input_ids,
causal_mask=causal_mask,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
pass
hidden_states = outputs[0]
bsz, q_len, hd = hidden_states.shape