From 82ead808d59958e767fdc22e245916f529e2b4cd Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Fri, 2 Feb 2024 20:19:15 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 82 ++++++++++++++++++++++++++++++++++------- 1 file changed, 68 insertions(+), 14 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 6812b2cba9..efe34e44df 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -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