diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index ccd9f89948..a341f3e154 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -18,6 +18,7 @@ import torch from .utils import calculate_settings +@triton.heuristics({"ADD_ONE": lambda args: args["ADD_ONE"],}) @triton.jit def _rms_layernorm_forward( Y, Y_row_stride, @@ -25,7 +26,8 @@ def _rms_layernorm_forward( W, W_row_stride, r, r_row_stride, n_cols, eps, - BLOCK_SIZE : tl.constexpr + ADD_ONE: tl.constexpr, + BLOCK_SIZE : tl.constexpr, ): """ Fast RMS Layernorm kernel @@ -48,11 +50,16 @@ def _rms_layernorm_forward( tl.store(r, inv_var) normed = X_row * inv_var normed = normed.to(W_row.dtype) # Exact copy from HF - output = normed * W_row + + # For Gemma - cannot do += 1 since float16 - maybe use FMADD + if not ADD_ONE: output = normed * W_row + else: output = normed * (W_row + 1.0) + tl.store(Y + col_offsets, output, mask = mask) pass +@triton.heuristics({"ADD_ONE": lambda args: args["ADD_ONE"],}) @triton.jit def _rms_layernorm_backward( dY, dY_row_stride, @@ -61,6 +68,7 @@ def _rms_layernorm_backward( r, r_row_stride, dW, dW_row_stride, n_cols, eps, + ADD_ONE: tl.constexpr, BLOCK_SIZE : tl.constexpr, ): """ @@ -84,7 +92,10 @@ def _rms_layernorm_backward( inv_var = tl.load(r).to(tl.float32) normed = X_row * inv_var - dY_W = dY_row * W_row + # For Gemma - cannot do += 1 since float16 - maybe use FMADD + if not ADD_ONE: dY_W = dY_row * W_row + else: dY_W = dY_row * (W_row + 1.0) + rowsum_dY_normed = tl.sum(dY_W * normed, axis = 0) output = inv_var/n_cols * (n_cols*dY_W - normed*rowsum_dY_normed) tl.store(dY + col_offsets, output, mask = mask) @@ -93,7 +104,7 @@ pass class Fast_RMS_Layernorm(torch.autograd.Function): @staticmethod - def forward(ctx, X, W, eps): + def forward(ctx, X, W, eps, add_one = False): shape = X.shape dim = shape[-1] X = X.view(-1, dim) @@ -109,12 +120,14 @@ class Fast_RMS_Layernorm(torch.autograd.Function): W, W.stride(0), r, r.stride(0), n_cols, eps, + ADD_ONE = add_one, BLOCK_SIZE = BLOCK_SIZE, num_warps = num_warps, ) ctx.eps = eps ctx.BLOCK_SIZE = BLOCK_SIZE ctx.num_warps = num_warps + ctx.ADD_ONE = add_one ctx.save_for_backward(X, W, r) return Y.view(*shape) pass @@ -135,18 +148,19 @@ class Fast_RMS_Layernorm(torch.autograd.Function): r, r .stride(0), dW, dW.stride(0), n_cols, ctx.eps, + ADD_ONE = ctx.ADD_ONE, BLOCK_SIZE = ctx.BLOCK_SIZE, num_warps = ctx.num_warps, ) dX = dY.view(*shape) - return dX, None, None + return dX, None, None, None pass pass -def fast_rms_layernorm(layernorm, X): +def fast_rms_layernorm(layernorm, X, add_one = False): W = layernorm.weight eps = layernorm.variance_epsilon - out = Fast_RMS_Layernorm.apply(X, W, eps) + out = Fast_RMS_Layernorm.apply(X, W, eps, add_one = add_one) return out pass diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index c99bd5f927..92e6978b89 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -39,6 +39,7 @@ except: pass +torch_nn_functional_gelu = torch.nn.functional.gelu def fast_geglu_inference(self, X): # gate = self.gate_proj(X) # up = self.up_proj(X) @@ -48,7 +49,7 @@ def fast_geglu_inference(self, X): gate = fast_linear_forward(self.gate_proj, X, out = temp[0]) up = fast_linear_forward(self. up_proj, X, out = temp[1]) - gate = torch.nn.functional.gelu(gate, approximate = "tanh") + gate = torch_nn_functional_gelu(gate, approximate = "tanh") gate *= up # X = self.down_proj(gate) @@ -57,6 +58,19 @@ def fast_geglu_inference(self, X): pass +def fast_rms_layernorm_inference_add_one(self, X, out_weight = None): + old_dtype = X.dtype + 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 + out_weight = torch.add(self.weight, 1.0, out = out_weight) + X *= out_weight + return XXX +pass + + # https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L590 def GemmaDecoderLayer_fast_forward( self, @@ -75,7 +89,7 @@ def GemmaDecoderLayer_fast_forward( # Self Attention residual = hidden_states - hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states) + hidden_states = fast_rms_layernorm_inference_add_one(self.input_layernorm, hidden_states) hidden_states, present_key_value = LlamaAttention_fast_forward_inference( self.self_attn, hidden_states, @@ -87,12 +101,12 @@ def GemmaDecoderLayer_fast_forward( # Fully Connected residual = hidden_states - hidden_states = fast_rms_layernorm_inference(self.post_attention_layernorm, hidden_states) + hidden_states = fast_rms_layernorm_inference_add_one(self.post_attention_layernorm, hidden_states) hidden_states = fast_geglu_inference(self.mlp, hidden_states) hidden_states += residual else: residual = hidden_states - hidden_states = fast_rms_layernorm(self.input_layernorm, hidden_states) + hidden_states = fast_rms_layernorm(self.input_layernorm, hidden_states, add_one = True) # hidden_states = self.input_layernorm(hidden_states) hidden_states, self_attn_weights, present_key_value = self.self_attn( hidden_states=hidden_states, @@ -108,7 +122,7 @@ def GemmaDecoderLayer_fast_forward( # Fully Connected residual = hidden_states - hidden_states = fast_rms_layernorm(self.post_attention_layernorm, hidden_states) + hidden_states = fast_rms_layernorm(self.post_attention_layernorm, hidden_states, add_one = True) # hidden_states = self.post_attention_layernorm(hidden_states) hidden_states = self.mlp(hidden_states) hidden_states = residual + hidden_states @@ -137,15 +151,16 @@ def GemmaModel_fast_forward_inference( ): # Fix out of bounds tokenization input_ids = input_ids[:,:self.max_seq_length] + out_weight = torch.empty_like(self.layers[0].input_layernorm.weight) hidden_states = self.embed_tokens(input_ids) hidden_states *= math_sqrt(self.config.hidden_size) - + 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 = fast_rms_layernorm_inference_add_one(decoder_layer.input_layernorm, hidden_states, out_weight) hidden_states, present_key_value = LlamaAttention_fast_forward_inference( decoder_layer.self_attn, hidden_states, @@ -156,13 +171,13 @@ def GemmaModel_fast_forward_inference( # Fully Connected residual = hidden_states - hidden_states = fast_rms_layernorm_inference(decoder_layer.post_attention_layernorm, hidden_states) + hidden_states = fast_rms_layernorm_inference_add_one(decoder_layer.post_attention_layernorm, hidden_states, out_weight) hidden_states = fast_geglu_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) + hidden_states = fast_rms_layernorm_inference_add_one(self.norm, hidden_states, out_weight) return BaseModelOutputWithPast( last_hidden_state = hidden_states, @@ -173,94 +188,6 @@ def GemmaModel_fast_forward_inference( pass -def GemmaForCausalLM_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]: - - if causal_mask is None and past_key_values is None: - causal_mask = xformers.attn_bias.LowerTriangularMask() - - 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) - self.model._has_no_labels = labels is None - - if past_key_values is not None and \ - hasattr(self.model.layers[0].self_attn, "paged_attention"): - outputs = GemmaModel_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 - if bsz == 1 and q_len == 1: - logits = torch.mv(self.lm_head.weight, hidden_states.ravel()) - logits = logits.unsqueeze(0).unsqueeze(0) - else: - logits = self.lm_head(hidden_states) - pass - - loss = None - if labels is not None: - shift_logits = logits - if not hasattr(self, "extra_ignored_labels"): - # Fixes https://github.com/unslothai/unsloth/issues/10 - self.extra_ignored_labels = torch.full((self.max_seq_length, 1), -100, device = "cuda") - pass - - shift_labels = torch.hstack((labels[..., 1:], self.extra_ignored_labels[:labels.shape[0]])) - loss = fast_cross_entropy_loss( - logits = shift_logits, - labels = shift_labels, - ) - pass - - 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 - - # Follows line by line https://github.com/google-deepmind/gemma/blob/main/gemma/positional_embeddings.py#L45 # Formulates cos and sin differently from Llama! class GemmaFixedRotaryEmbedding(torch.nn.Module): @@ -290,7 +217,7 @@ class GemmaFixedRotaryEmbedding(torch.nn.Module): positions = torch.arange(self.max_seq_len_cached, device = "cpu", dtype = torch.int64).float() radians_new = positions[..., None] / timescale[None, None, :] radians_new = radians_new.squeeze(0) - + emb = torch.cat((radians_new, radians_new), dim=-1) self.register_buffer("cos_cached", emb.cos().to(dtype=dtype, device=device, non_blocking=True), persistent=False) self.register_buffer("sin_cached", emb.sin().to(dtype=dtype, device=device, non_blocking=True), persistent=False) @@ -318,7 +245,7 @@ class FastGemmaModel(FastLlamaModel): GemmaFlashAttention2.forward = LlamaAttention_fast_forward GemmaDecoderLayer .forward = GemmaDecoderLayer_fast_forward GemmaModel .forward = LlamaModel_fast_forward - GemmaForCausalLM .forward = GemmaForCausalLM_fast_forward + GemmaForCausalLM .forward = CausalLM_fast_forward(GemmaModel_fast_forward_inference) PeftModelForCausalLM.forward = PeftModelForCausalLM_fast_forward # Solves https://github.com/unslothai/unsloth/issues/168 # Static KV Cache was introduced in 4.38.0, causing training to be much slower. @@ -406,8 +333,8 @@ class FastGemmaModel(FastLlamaModel): # Must be in float32 # https://github.com/keras-team/keras-nlp/blob/v0.8.2/keras_nlp/models/gemma/rms_normalization.py#L36 # module = module.to(torch.float32) - # Don't convert to float32 since error analysis shows it makes it worse!! - module.weight += 1.0 # return output * (1 + self.weight) + # Leave + 1 to Triton kernel itself + # module.weight += 1.0 # return output * (1 + self.weight) if not hasattr(module, "variance_epsilon"): module.variance_epsilon = module.eps # Gemma doesn't use variance_epsilon pass diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 20552644d9..05cd1c7426 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -208,6 +208,7 @@ def LlamaAttention_fast_forward_inference( pass +torch_nn_functional_silu = torch.nn.functional.silu def fast_swiglu_inference(self, X): # gate = self.gate_proj(X) # up = self.up_proj(X) @@ -217,7 +218,7 @@ def fast_swiglu_inference(self, X): gate = fast_linear_forward(self.gate_proj, X, out = temp[0]) up = fast_linear_forward(self. up_proj, X, out = temp[1]) - gate = torch.nn.functional.silu(gate, inplace = True) + gate = torch_nn_functional_silu(gate, inplace = True) gate *= up # X = self.down_proj(gate) @@ -509,7 +510,8 @@ def LlamaModel_fast_forward( inputs_embeds = self.embed_tokens(input_ids) # Mormalized from Gemma - if self.config.model_type == "gemma": + IS_GEMMA = self.config.model_type == "gemma" + if IS_GEMMA: inputs_requires_grad = inputs_embeds.requires_grad if not inputs_embeds.is_leaf: inputs_embeds = inputs_embeds.detach() @@ -619,7 +621,7 @@ def LlamaModel_fast_forward( all_self_attns += (layer_outputs[1],) pass - hidden_states = fast_rms_layernorm(self.norm, hidden_states) + hidden_states = fast_rms_layernorm(self.norm, hidden_states, add_one = IS_GEMMA) # add hidden states from the last decoder layer if output_hidden_states: @@ -681,91 +683,94 @@ def LlamaModel_fast_forward_inference( pass -def LlamaForCausalLM_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]: +def CausalLM_fast_forward(fast_forward_inference): + def _CausalLM_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]: - if causal_mask is None and past_key_values is None: - causal_mask = xformers.attn_bias.LowerTriangularMask() + if causal_mask is None and past_key_values is None: + causal_mask = xformers.attn_bias.LowerTriangularMask() - 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) - self.model._has_no_labels = labels is None - - if past_key_values 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, + 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 ) - 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 + return_dict = return_dict if return_dict is not None else self.config.use_return_dict - hidden_states = outputs[0] - bsz, q_len, hd = hidden_states.shape - if bsz == 1 and q_len == 1: - logits = torch.mv(self.lm_head.weight, hidden_states.ravel()) - logits = logits.unsqueeze(0).unsqueeze(0) - else: - logits = self.lm_head(hidden_states) - pass + # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) + self.model._has_no_labels = labels is None - loss = None - if labels is not None: - shift_logits = logits - if not hasattr(self, "extra_ignored_labels"): - # Fixes https://github.com/unslothai/unsloth/issues/10 - self.extra_ignored_labels = torch.full((self.max_seq_length, 1), -100, device = "cuda") + if past_key_values is not None and \ + hasattr(self.model.layers[0].self_attn, "paged_attention"): + outputs = 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 - - shift_labels = torch.hstack((labels[..., 1:], self.extra_ignored_labels[:labels.shape[0]])) - loss = fast_cross_entropy_loss( - logits = shift_logits, - labels = shift_labels, + + hidden_states = outputs[0] + bsz, q_len, hd = hidden_states.shape + if bsz == 1 and q_len == 1: + logits = torch.mv(self.lm_head.weight, hidden_states.ravel()) + logits = logits.unsqueeze(0).unsqueeze(0) + else: + logits = self.lm_head(hidden_states) + pass + + loss = None + if labels is not None: + shift_logits = logits + if not hasattr(self, "extra_ignored_labels"): + # Fixes https://github.com/unslothai/unsloth/issues/10 + self.extra_ignored_labels = torch.full((self.max_seq_length, 1), -100, device = "cuda") + pass + + shift_labels = torch.hstack((labels[..., 1:], self.extra_ignored_labels[:labels.shape[0]])) + loss = fast_cross_entropy_loss( + logits = shift_logits, + labels = shift_labels, + ) + pass + + 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 - - 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, - ) + return _CausalLM_fast_forward pass @@ -880,7 +885,7 @@ class FastLlamaModel: LlamaFlashAttention2.forward = LlamaAttention_fast_forward LlamaDecoderLayer .forward = LlamaDecoderLayer_fast_forward LlamaModel .forward = LlamaModel_fast_forward - LlamaForCausalLM .forward = LlamaForCausalLM_fast_forward + LlamaForCausalLM .forward = CausalLM_fast_forward(LlamaModel_fast_forward_inference) PeftModelForCausalLM.forward = PeftModelForCausalLM_fast_forward # Solves https://github.com/unslothai/unsloth/issues/168