From 64fcf4dbe7c47fe5506bfa0b784733179956ead0 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 8 Jan 2026 12:10:29 +0000 Subject: [PATCH] Fix Qwen3-Omni model loading and training support - Handle list model_type in loader.py dispatch (multimodal models return a list) - Add forward method to Qwen3OmniMoeForConditionalGeneration that delegates to thinker - Add embedding accessor methods (get_input_embeddings, set_input_embeddings, etc.) - Refactor FastQwen3OmniMoeModel.from_pretrained to use proper multimodal loading pattern - Use AutoProcessor instead of tokenizer for multimodal support The Qwen3OmniMoeForConditionalGeneration model only has a generate() method by default. Training requires a forward() method that computes loss, which now delegates to self.thinker.forward() where the actual causal LM forward pass is implemented. Tested with 30 training steps on Qwen/Qwen3-Omni-30B-A3B-Instruct: - Loss: 1.28 -> 0.81 (minimum at step 28) - Grad norms stable around 0.3-0.9 - Training completed successfully in ~14.5 minutes --- unsloth/models/loader.py | 7 +- unsloth/models/qwen3_moe.py | 285 +++++++++++++++++++++++++----------- 2 files changed, 204 insertions(+), 88 deletions(-) diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index 572dd43e74..ac13cbb994 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -619,6 +619,7 @@ class FastLanguageModel(FastLlamaModel): model_type == "qwen3" or model_type == "qwen3_moe" or model_type == "qwen3_omni_moe" + or (isinstance(model_type, (list, tuple)) and "qwen3_omni_moe" in model_type) ): if not SUPPORTS_QWEN3 or not SUPPORTS_QWEN3_MOE: raise ImportError( @@ -627,7 +628,11 @@ class FastLanguageModel(FastLlamaModel): f'Try `pip install --upgrade "transformers>=4.50.3"`\n' f"to obtain the latest transformers build, then restart this session." ) - if model_type == "qwen3_omni_moe": + # For multimodal models like Qwen3-Omni, model_type can be a list + is_qwen3_omni = model_type == "qwen3_omni_moe" or ( + isinstance(model_type, (list, tuple)) and "qwen3_omni_moe" in model_type + ) + if is_qwen3_omni: dispatch_model = FastQwen3OmniMoeModel else: dispatch_model = ( diff --git a/unsloth/models/qwen3_moe.py b/unsloth/models/qwen3_moe.py index f14f03dcb7..edff64523b 100644 --- a/unsloth/models/qwen3_moe.py +++ b/unsloth/models/qwen3_moe.py @@ -257,107 +257,218 @@ def qwen3_omni_apply_o(self, X): return O -class FastQwen3OmniMoeModel(FastQwen3MoeModel): - @staticmethod - def from_pretrained( - model_name = "Qwen/Qwen3-Omni-30B-A3B-Instruct", - max_seq_length = 4096, - dtype = None, - load_in_4bit = True, - token = None, - device_map = "sequential", - rope_scaling = None, - fix_tokenizer = True, - model_patcher = None, - tokenizer_name = None, - trust_remote_code = False, +def _patch_qwen3_omni_model_class(model_class): + """ + Patch Qwen3OmniMoeForConditionalGeneration with required methods for training. + The model delegates to its thinker submodule for actual forward pass. + """ + # Add embedding accessor methods that delegate to thinker + def get_input_embeddings(self): + return self.thinker.get_input_embeddings() + + def set_input_embeddings(self, value): + self.thinker.set_input_embeddings(value) + + def get_output_embeddings(self): + return self.thinker.lm_head + + def set_output_embeddings(self, value): + self.thinker.lm_head = value + + # Add forward method that delegates to thinker for training + def forward( + self, + input_ids=None, + input_features=None, + pixel_values=None, + pixel_values_videos=None, + image_grid_thw=None, + video_grid_thw=None, + attention_mask=None, + feature_attention_mask=None, + audio_feature_lengths=None, + position_ids=None, + past_key_values=None, + inputs_embeds=None, + rope_deltas=None, + labels=None, + use_cache=None, + output_router_logits=None, + use_audio_in_video=None, + cache_position=None, + video_second_per_grid=None, **kwargs, ): - model, tokenizer = FastQwen3MoeModel.from_pretrained( - model_name = model_name, - max_seq_length = max_seq_length, - dtype = dtype, - load_in_4bit = load_in_4bit, - token = token, - device_map = device_map, - rope_scaling = rope_scaling, - fix_tokenizer = fix_tokenizer, - model_patcher = FastQwen3Model, - tokenizer_name = tokenizer_name, - trust_remote_code = trust_remote_code, + """Forward method that delegates to the thinker for training.""" + return self.thinker.forward( + input_ids=input_ids, + input_features=input_features, + pixel_values=pixel_values, + pixel_values_videos=pixel_values_videos, + image_grid_thw=image_grid_thw, + video_grid_thw=video_grid_thw, + attention_mask=attention_mask, + feature_attention_mask=feature_attention_mask, + audio_feature_lengths=audio_feature_lengths, + position_ids=position_ids, + past_key_values=past_key_values, + inputs_embeds=inputs_embeds, + rope_deltas=rope_deltas, + labels=labels, + use_cache=use_cache, + output_router_logits=output_router_logits, + use_audio_in_video=use_audio_in_video, + cache_position=cache_position, + video_second_per_grid=video_second_per_grid, **kwargs, ) - # Manually patch Thinker, Talker, and MTP modules - patched_classes = set() + # Apply patches to the model class + if not hasattr(model_class, "_unsloth_patched"): + model_class.get_input_embeddings = get_input_embeddings + model_class.set_input_embeddings = set_input_embeddings + model_class.get_output_embeddings = get_output_embeddings + model_class.set_output_embeddings = set_output_embeddings + model_class.forward = forward + model_class._unsloth_patched = True - for name, module in model.named_modules(): - if module.__class__ in patched_classes: - continue - class_name = module.__class__.__name__ +def _patch_qwen3_omni_attention_modules(model): + """ + Patch Attention and MLP modules in Qwen3-Omni for faster training. + """ + patched_classes = set() - # Replace Attention with Triton Kernels + for name, module in model.named_modules(): + if module.__class__ in patched_classes: + continue + + class_name = module.__class__.__name__ + + # Replace Attention with Triton Kernels + if ( + "Attention" in class_name + and hasattr(module, "q_norm") + and hasattr(module, "k_norm") + ): if ( - "Attention" in class_name - and hasattr(module, "q_norm") - and hasattr(module, "k_norm") + hasattr(module, "q_proj") + and hasattr(module, "k_proj") + and hasattr(module, "v_proj") + and hasattr(module, "o_proj") ): - if ( - hasattr(module, "q_proj") - and hasattr(module, "k_proj") - and hasattr(module, "v_proj") - and hasattr(module, "o_proj") + module.__class__.apply_qkv = qwen3_omni_apply_qkv + module.__class__.apply_o = qwen3_omni_apply_o + + # Save original forward to allow fallback during inference + if not hasattr(module.__class__, "_original_forward"): + module.__class__._original_forward = module.__class__.forward + + # Define a Safe Wrapper + def _attention_wrapper( + self, + hidden_states, + position_embeddings=None, + attention_mask=None, + past_key_values=None, + **kwargs, ): - module.__class__.apply_qkv = qwen3_omni_apply_qkv - module.__class__.apply_o = qwen3_omni_apply_o - - # Save original forward to allow fallback during inference - if not hasattr(module.__class__, "_original_forward"): - module.__class__._original_forward = module.__class__.forward - - # Define a Safe Wrapper - def _attention_wrapper( - self, - hidden_states, - position_embeddings = None, - attention_mask = None, - past_key_values = None, - **kwargs, - ): - # FALLBACK: If caching is used (Inference), use original code to avoid breaking Cache objects - if past_key_values is not None: - return self._original_forward( - hidden_states, - position_embeddings = position_embeddings, - attention_mask = attention_mask, - past_key_values = past_key_values, - **kwargs, - ) - - # If training (No Cache), use Unsloth's 2x Faster Kernel - return Qwen3Attention_fast_forward( - self, + # FALLBACK: If caching is used (Inference), use original code + if past_key_values is not None: + return self._original_forward( hidden_states, - attention_mask = attention_mask, - position_embeddings = position_embeddings, - past_key_value = None, # Explicitly None for training + position_embeddings=position_embeddings, + attention_mask=attention_mask, + past_key_values=past_key_values, **kwargs, ) - # Apply the wrapper - module.__class__.forward = _attention_wrapper - patched_classes.add(module.__class__) + # If training (No Cache), use Unsloth's 2x Faster Kernel + return Qwen3Attention_fast_forward( + self, + hidden_states, + attention_mask=attention_mask, + position_embeddings=position_embeddings, + past_key_value=None, + **kwargs, + ) - # Replace MLP with SwiGLU Kernels - # Targets: Qwen3OmniMoeThinkerTextMLP, Qwen3OmniMoeCode2WavMlp - if "MLP" in class_name and "Moe" not in class_name: - if ( - hasattr(module, "gate_proj") - and hasattr(module, "up_proj") - and hasattr(module, "down_proj") - ): - module.__class__.forward = fast_swiglu_inference - patched_classes.add(module.__class__) + module.__class__.forward = _attention_wrapper + patched_classes.add(module.__class__) - return model, tokenizer + # Replace MLP with SwiGLU Kernels + if "MLP" in class_name and "Moe" not in class_name: + if ( + hasattr(module, "gate_proj") + and hasattr(module, "up_proj") + and hasattr(module, "down_proj") + ): + module.__class__.forward = fast_swiglu_inference + patched_classes.add(module.__class__) + + return patched_classes + + +class FastQwen3OmniMoeModel(FastQwen3MoeModel): + @staticmethod + def from_pretrained( + model_name="Qwen/Qwen3-Omni-30B-A3B-Instruct", + max_seq_length=4096, + dtype=None, + load_in_4bit=True, + token=None, + device_map="sequential", + rope_scaling=None, + fix_tokenizer=True, + model_patcher=None, + tokenizer_name=None, + trust_remote_code=False, + **kwargs, + ): + # Import the model class + from transformers import Qwen3OmniMoeForConditionalGeneration, AutoProcessor + + # Patch the model class with required methods before loading + _patch_qwen3_omni_model_class(Qwen3OmniMoeForConditionalGeneration) + + # Determine dtype + if dtype is None: + dtype = torch.bfloat16 + + # Load the model using the multimodal pattern + model_kwargs = { + "torch_dtype": dtype, + "trust_remote_code": trust_remote_code, + "token": token, + "device_map": device_map, + } + + if load_in_4bit: + from transformers import BitsAndBytesConfig + model_kwargs["quantization_config"] = BitsAndBytesConfig( + load_in_4bit=True, + bnb_4bit_compute_dtype=dtype, + bnb_4bit_use_double_quant=True, + bnb_4bit_quant_type="nf4", + ) + + # Merge any additional kwargs + model_kwargs.update(kwargs) + + # Load model + model = Qwen3OmniMoeForConditionalGeneration.from_pretrained( + model_name, + **model_kwargs, + ) + + # Load processor/tokenizer + processor = AutoProcessor.from_pretrained( + model_name, + trust_remote_code=trust_remote_code, + token=token, + ) + + # Patch attention and MLP modules for faster training + patched_classes = _patch_qwen3_omni_attention_modules(model) + + return model, processor