diff --git a/unsloth/models/phi2.py b/unsloth/models/phi2.py index 0949b9bd19..1e8f076ba2 100644 --- a/unsloth/models/phi2.py +++ b/unsloth/models/phi2.py @@ -16,6 +16,8 @@ from .llama import * from ._utils import __version__ from ..kernels.relu import relu_kernel +from torch.nn import CrossEntropyLoss + from transformers.models.phi.modeling_phi import ( PhiAttention, PhiDecoderLayer, @@ -155,7 +157,74 @@ def Phi2Attention_fast_forward( return attn_output, attn_weights, past_key_value pass -inplace_rope_embedding + +def Phi2ForCausalLM_fast_forward( + self, + input_ids: torch.LongTensor = None, + causal_mask: Optional[xformers.attn_bias.BlockDiagonalCausalMask] = None, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.LongTensor] = None, + past_key_values: Optional[List[torch.FloatTensor]] = None, + inputs_embeds: Optional[torch.FloatTensor] = None, + labels: Optional[torch.LongTensor] = None, + use_cache: Optional[bool] = None, + output_attentions: Optional[bool] = None, + output_hidden_states: Optional[bool] = None, + return_dict: Optional[bool] = None, + *args, **kwargs, +) -> Union[Tuple, CausalLMOutputWithPast]: + + output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) + outputs = self.model( + input_ids=input_ids, + 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, + ) + + hidden_states = outputs[0] + logits = self.lm_head(hidden_states) + logits = logits.float() + + loss = None + if labels is not None: + # Shift so that tokens < n predict n + shift_logits = logits[..., :-1, :].contiguous() + shift_labels = labels[..., 1:].contiguous() + # Flatten the tokens + shift_logits = shift_logits.view(-1, self.config.vocab_size) + shift_labels = shift_labels.view(-1) + # Enable model parallelism + shift_labels = shift_labels.to(shift_logits.device) + loss = fast_cross_entropy_loss( + logits = shift_logits, + labels = shift_labels, + ) + if not return_dict: + output = (logits,) + outputs[1:] + return (loss,) + output if loss is not None else output + + return CausalLMOutputWithPast( + loss=loss, + logits=logits, + past_key_values=outputs.past_key_values, + hidden_states=outputs.hidden_states, + attentions=outputs.attentions, + ) +pass + + def fast_mlp_inference(self, X): gate = self.gate_proj(X) up = self.up_proj(X) @@ -171,6 +240,11 @@ class FastPhi2Model(FastLlamaModel): def pre_patch(): PhiAttention .forward = Phi2Attention_fast_forward PhiFlashAttention2 .forward = Phi2Attention_fast_forward + PhiDecoderLayer .forward = LlamaDecoderLayer_fast_forward + PhiModel .forward = LlamaModel_fast_forward + PhiForCausalLM .forward = Phi2ForCausalLM_fast_forward + PeftModelForCausalLM.forward = PeftModelForCausalLM_fast_forward + pass @staticmethod