From ea9b4eea0cb07bf02339f226d285e76e3b4d94e0 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Fri, 2 Feb 2024 23:40:52 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 61 ++++++++++++----------------------------- 1 file changed, 17 insertions(+), 44 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index ce21733056..e7cffbb7b6 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -121,7 +121,7 @@ def LlamaAttention_fast_forward_inference_prefill( # Prefill phase # if not hasattr(self, "paged_attention"): - self.paged_attention = torch.empty((self.max_seq_length+1, 2, bsz, n_kv_heads, head_dim), dtype = dtype, device = "cuda") + self.paged_attention = torch.empty((self.config.max_position_embeddings+1, 2, bsz, n_kv_heads, head_dim), dtype = dtype, device = "cuda") self.paged_attention_K = self.paged_attention[:,0] self.paged_attention_V = self.paged_attention[:,1] self.paged_attention_K[:seq_len] = K1.permute(2, 0, 1, 3) @@ -129,7 +129,7 @@ def LlamaAttention_fast_forward_inference_prefill( self.temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda") self.temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda") self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda") - self.attention = torch.empty((bsz, n_heads, 1, self.max_seq_length), dtype = dtype, device = "cuda") + self.attention = torch.empty((bsz, n_heads, 1, self.config.max_position_embeddings), dtype = dtype, device = "cuda") self.scalar = 1.0 / math_sqrt(self.head_dim) # pass @@ -252,8 +252,12 @@ pass def fast_mlp_inference(self, X): # gate = self.gate_proj(X) # up = self.up_proj(X) - gate = fast_linear_forward(self.gate_proj, X, out = self.temp_buffer[0]) - up = fast_linear_forward(self. up_proj, X, out = self.temp_buffer[1]) + bsz, _, hd = X.shape + mlp_size = self.config.intermediate_size + temp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda") + + 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 *= up @@ -265,20 +269,11 @@ pass def fast_rms_layernorm_inference(self, X): old_dtype = X.dtype - XX = self.temp_buffer1[0] - XX[:] = X # XX = X.to(torch.float32) - - # variance = XX.square().mean(-1, keepdim = True) - torch.square(XX, out = self.temp_buffer1[1]) - variance = self.temp_buffer1[1].mean(-1, keepdim = True) - + 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 - self.temp_buffer2[:] = XX - X = self.temp_buffer2 - + X = XX.to(old_dtype) # Must preserve due to residual X *= self.weight return X pass @@ -421,22 +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"): - - bsz, _, hd = hidden_states.shape - dtype = hidden_states.dtype - - # Create temp matrices for layernorms - temp_buffer1 = torch.empty((2, bsz, 1, hd), dtype = torch.float32, device = "cuda") - temp_buffer2 = torch.empty((bsz, 1, hd), dtype = dtype, device = "cuda") - self.input_layernorm.temp_buffer1 = temp_buffer1 - self.input_layernorm.temp_buffer2 = temp_buffer2 - self.post_attention_layernorm.temp_buffer1 = temp_buffer1 - self.post_attention_layernorm.temp_buffer2 = temp_buffer2 - - # Create temp matrices for MLP - mlp_size = self.config.intermediate_size - self.mlp.temp_buffer = torch.empty((2, bsz, 1, mlp_size), dtype = dtype, device = "cuda") - # Self Attention residual = hidden_states hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states) @@ -453,17 +432,7 @@ def LlamaDecoderLayer_fast_forward( hidden_states = fast_rms_layernorm_inference(self.post_attention_layernorm, hidden_states) hidden_states = fast_mlp_inference(self.mlp, hidden_states) hidden_states += residual - elif past_key_value is not None: - # Delete old temporary buffers - if hasattr(self.input_layernorm, "temp_buffer1"): - del self.post_attention_layernorm.temp_buffer1 - del self.post_attention_layernorm.temp_buffer2 - del self.post_attention_layernorm.temp_buffer1 - del self.post_attention_layernorm.temp_buffer2 - del self.mlp.temp_buffer - pass - # Self Attention residual = hidden_states hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states) @@ -687,7 +656,11 @@ def LlamaModel_fast_forward( all_self_attns += (layer_outputs[1],) pass - hidden_states = fast_rms_layernorm(self.norm, hidden_states) + if past_key_values is not None: + hidden_states = fast_rms_layernorm_inference(self.norm, hidden_states) + else: + hidden_states = fast_rms_layernorm(self.norm, hidden_states) + pass # add hidden states from the last decoder layer if output_hidden_states: @@ -706,7 +679,7 @@ pass # https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L825 -@torch.inference_mode +@torch.compile def LlamaModel_fast_forward_inference( self, input_ids,