From eb495f171f8df7d55c74721eab548a97377539cf Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Fri, 2 May 2025 03:09:44 -0700 Subject: [PATCH] Qwen 3 --- unsloth/models/llama4.py | 4 +- unsloth/models/qwen3.py | 99 ++-------------------------------------- 2 files changed, 5 insertions(+), 98 deletions(-) diff --git a/unsloth/models/llama4.py b/unsloth/models/llama4.py index 9818b3db04..f429922268 100644 --- a/unsloth/models/llama4.py +++ b/unsloth/models/llama4.py @@ -12,5 +12,5 @@ # See the License for the specific language governing permissions and # limitations under the License. -from unsloth_studio.models import patch_llama4 -patch_llama4() +# from unsloth_studio.models import patch_llama4 +# patch_llama4() diff --git a/unsloth/models/qwen3.py b/unsloth/models/qwen3.py index c0ceefd10c..ac3ac108fb 100644 --- a/unsloth/models/qwen3.py +++ b/unsloth/models/qwen3.py @@ -15,6 +15,7 @@ from .llama import * import os from ._utils import __version__ +from unsloth_zoo.utils import Version, _get_dtype from .llama import ( LlamaRotaryEmbedding, LlamaLinearScalingRotaryEmbedding, @@ -28,7 +29,6 @@ try: Qwen3ForCausalLM, ) except: - from packaging.version import Version transformers_version = Version(transformers_version) if not transformers_version >= Version("4.50.3"): #TODO: Update when transformers is updated raise ImportError( @@ -51,7 +51,6 @@ except: Qwen3SdpaAttention = Qwen3Attention Qwen3FlashAttention2 = Qwen3Attention pass -from unsloth_zoo.utils import Version, _get_dtype def Qwen3Attention_fast_forward( @@ -286,8 +285,8 @@ def Qwen3Attention_fast_forward_inference( Kn = Kn.view(bsz, 1, n_kv_heads, head_dim)#.transpose(1, 2) # we will transpose after normalisation Vn = Vn.view(bsz, 1, n_kv_heads, head_dim).transpose(1, 2) - Qn = fast_rms_layernorm(self.q_norm, Qn) - Kn = fast_rms_layernorm(self.k_norm, Kn) + Qn = fast_rms_layernorm_inference(self.q_norm, Qn) + Kn = fast_rms_layernorm_inference(self.k_norm, Kn) Qn = Qn.transpose(1, 2) Kn = Kn.transpose(1, 2) @@ -368,98 +367,6 @@ def Qwen3Attention_fast_forward_inference( return A, (Kn, Vn) pass -# def Qwen3Model_fast_forward_inference( -# self, -# input_ids, -# past_key_values, -# position_ids, -# attention_mask = None, -# ): -# input_ids = input_ids[:,:self.max_seq_length] -# bsz, q_len = input_ids.shape -# hd = self.config.hidden_size -# mlp_size = self.config.intermediate_size - -# X = self.model.embed_tokens(input_ids) -# X = X.to(_get_dtype(self.config.torch_dtype)) -# bsz, q_len, hd = X.shape -# assert(q_len == 1) -# # Get saved buffers to reduce memory movement -# residual = torch.empty((bsz, q_len, hd), dtype = torch.float32, device = "cuda:0") -# _XX = torch.empty((2, bsz, q_len, hd), dtype = torch.float32, device = "cuda:0") -# XX, XX2 = _XX[0], _XX[1] -# variance = torch.empty((bsz, q_len, 1), dtype = torch.float32, device = "cuda:0") -# temp_mlp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda:0") -# temp_gate, temp_up = temp_mlp[0], temp_mlp[1] - -# seq_len = past_key_values[0][0].shape[-2] -# if bsz != 1: -# attention_mask = _prepare_4d_causal_attention_mask_for_sdpa( -# attention_mask, -# (bsz, q_len), -# X, -# seq_len, -# sliding_window = getattr(self.config, "sliding_window", None), -# ) -# else: -# attention_mask = None -# pass - -# next_decoder_cache = [] - -# for idx, decoder_layer in enumerate(self.model.layers): -# residual.copy_(X) # residual = X -# X = fast_rms_layernorm_inference( -# decoder_layer.input_layernorm, -# X, -# XX = XX, -# XX2 = XX2, -# variance = variance, -# ) -# X, present_key_value = Qwen3Attention_fast_forward_inference( -# decoder_layer.self_attn, -# hidden_states = X, -# past_key_value = past_key_values[idx], -# position_ids = position_ids, -# attention_mask = attention_mask, -# do_prefill = not hasattr(decoder_layer.self_attn, "paged_attention"), -# ) -# X += residual - -# residual.copy_(X) # residual = X -# X = fast_rms_layernorm_inference( -# decoder_layer.post_attention_layernorm, -# X, -# XX = XX, -# XX2 = XX2, -# variance = variance, -# ) -# X = fast_swiglu_inference( -# decoder_layer.mlp, -# X, -# temp_gate = temp_gate, -# temp_up = temp_up, -# ) -# X += residual - -# next_decoder_cache.append(present_key_value) -# pass -# X = fast_rms_layernorm_inference( -# self.model.norm, -# X, -# XX = XX, -# XX2 = XX2, -# variance = variance, -# ) - -# return BaseModelOutputWithPast( -# last_hidden_state = X, -# past_key_values = next_decoder_cache, -# hidden_states = [], -# attentions = [], -# ) -# pass - class FastQwen3Model(FastLlamaModel): @staticmethod