From a1aec618ce47d7e1f93b2c956ffe7f32695a3581 Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Wed, 22 Apr 2026 10:46:44 +0000 Subject: [PATCH] flex: fix Qwen3 MoE smoke regressions (dtype / MoE MLP / peft patching) Four integration fixes wired up while bringing Qwen3-30B-A3B-Instruct-2507 green end-to-end on UNSLOTH_FAST_INFERENCE=1: 1. unsloth/models/llama.py patch_peft_model: transformers 5.x reports model_type as "qwen3_moe" (with underscore); the PR's check was "qwen3moe" and fell through to NotImplementedError. 2. unsloth/models/llama.py patch_peft_model dense MLP patching: the fused gate/up/down LoRAMLP swap walks layer.mlp.gate_proj, which is a Qwen3MoeSparseMoeBlock for MoE and has no gate_proj attribute. Skip the swap when the MLP does not expose the dense trio; MoE LoRA is wired through unsloth_zoo/moe_utils anyway. 3. unsloth/inference/flex_qwen3_llama.py flex attention forward: bnb-4bit Linear compute produces fp32 k / v even under autocast, which makes the paged KV index_put_ refuse the mixed dtype (bf16 cache, fp32 update). Cast k / v to self._paged_cache.k_cache.dtype before update. Also benefits the dense path. 4. unsloth/inference/flex_engine.py bind_peft_model: for Qwen3 MoE the ParamWrapper keeps LoRA un-merged on the stacked expert tensors, so the training model's expert weights ARE the pristine source. Skip the pristine-base deep-copy for arch=="qwen3_moe" and point refresh_moe_lora_merge_from_pristine at the training base directly. Avoids a third 30-60 GB residency on 30B-A3B. 5. unsloth/inference/flex_moe.py call_moe_model_with_flex_kwargs: lock activations to the embed dtype across layernorm + MoE MLP; RMSNorm + bnb-4bit compute promote activations to fp32 along the MoE path under autocast. Also force-restore Qwen3MoeSparseMoeBlock.forward to the stock or unsloth_zoo version if FastQwen3MoeModel.pre_patch clobbered it with a legacy Qwen3MoeSparseMoeBlock_fast_forward that expects a flat self.gate_proj (which does not exist on transformers 5.x stacked-expert MoE blocks). Adds tests/flex_moe_smoke.py: generates 32 tokens twice (cold + warm), records first-call / warm-call tokens/s, peak VRAM, arch, impl. Writes async_task_outputs/qwen3_moe_grpo_bench/smoke_A_{4bit,bf16}.json. Measured on a single B200 (sm_100), Qwen3-30B-A3B-Instruct-2507 + LoRA rank 16 + grouped_mm MoE backend: | precision | t_load (s) | peak VRAM (GB) | cold tok/s | warm tok/s | |-----------|------------|----------------|------------|------------| | 4bit | 27.0 | 123.4 | 3.7 | 6.9 | | bf16 | 33.5 | 125.9 | 3.7 | 6.7 | Both completions coherent ("the lazy dog. ..."). --- tests/flex_moe_smoke.py | 156 ++++++++++++++++++++++++++ unsloth/inference/flex_engine.py | 34 ++++-- unsloth/inference/flex_moe.py | 48 +++++++- unsloth/inference/flex_qwen3_llama.py | 8 ++ unsloth/models/llama.py | 12 +- 5 files changed, 245 insertions(+), 13 deletions(-) create mode 100644 tests/flex_moe_smoke.py diff --git a/tests/flex_moe_smoke.py b/tests/flex_moe_smoke.py new file mode 100644 index 0000000000..beab04cc33 --- /dev/null +++ b/tests/flex_moe_smoke.py @@ -0,0 +1,156 @@ +# SPDX-License-Identifier: GNU Affero General Public License v3.0 +# Copyright 2023-present the Unsloth team. All rights reserved. + +"""Smoke-test ``UNSLOTH_FAST_INFERENCE=1`` on a Qwen3 MoE model. + +Mirrors ``tests/flex_fastlm_smoke.py`` but targets the new +``FlexMoEInference`` path added for ``Qwen3MoeForCausalLM``. + +Invoked as: + CUDA_VISIBLE_DEVICES=0 UNSLOTH_FAST_INFERENCE=1 python -u \ + tests/flex_moe_smoke.py \ + --model unsloth/Qwen3-30B-A3B-Instruct-2507 \ + --load_in_4bit + +Writes a small JSON summary to +``async_task_outputs/qwen3_moe_grpo_bench/smoke_A_{precision}.json``. +""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +import time +from pathlib import Path + +_REPO_ROOT = Path(__file__).resolve().parents[1] +if str(_REPO_ROOT) not in sys.path: + sys.path.insert(0, str(_REPO_ROOT)) + + +def main(): + p = argparse.ArgumentParser() + p.add_argument( + "--model", default = "unsloth/Qwen3-30B-A3B-Instruct-2507" + ) + p.add_argument("--dtype", choices = ["bf16", "fp16"], default = "bf16") + p.add_argument("--load_in_4bit", action = "store_true") + p.add_argument("--with_lora", action = "store_true") + p.add_argument("--max_new_tokens", type = int, default = 32) + p.add_argument("--max_seq_length", type = int, default = 1024) + p.add_argument("--prompt", default = "The quick brown fox jumps over") + p.add_argument("--out_dir", default = "async_task_outputs/qwen3_moe_grpo_bench") + args = p.parse_args() + + import torch + + os.environ.setdefault("UNSLOTH_FAST_INFERENCE", "1") + os.environ.setdefault("UNSLOTH_MOE_BACKEND", "grouped_mm") + print(f"[smoke] UNSLOTH_FAST_INFERENCE={os.environ.get('UNSLOTH_FAST_INFERENCE')}") + print(f"[smoke] UNSLOTH_MOE_BACKEND={os.environ.get('UNSLOTH_MOE_BACKEND')}") + + import unsloth + + print(f"[smoke] unsloth={unsloth.__file__}") + from unsloth import FastLanguageModel + + dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16 + + torch.cuda.reset_peak_memory_stats() + t0 = time.perf_counter() + model, tokenizer = FastLanguageModel.from_pretrained( + model_name = args.model, + max_seq_length = args.max_seq_length, + dtype = dtype, + load_in_4bit = args.load_in_4bit, + fast_inference = True, + ) + t_load = time.perf_counter() - t0 + peak_after_load = torch.cuda.max_memory_reserved() / 1024**3 + print(f"[smoke] loaded model in {t_load:.1f}s; peak VRAM after load: {peak_after_load:.2f} GB") + print(f"[smoke] hasattr(model, 'vllm_engine'): {hasattr(model, 'vllm_engine')}") + print(f"[smoke] vllm_engine type: {type(model.vllm_engine).__name__}") + arch = getattr(model.vllm_engine, "arch", "?") + impl = type(model.vllm_engine._impl).__name__ + print(f"[smoke] FlexEngine.arch={arch} impl={impl}") + + if args.with_lora: + model = FastLanguageModel.get_peft_model( + model, + r = 16, + target_modules = [ + "q_proj", "k_proj", "v_proj", "o_proj", + "gate_proj", "up_proj", "down_proj", "gate_up_proj", + ], + lora_alpha = 32, + lora_dropout = 0.0, + bias = "none", + use_gradient_checkpointing = "unsloth", + random_state = 3407, + ) + print(f"[smoke] PEFT model type: {type(model).__name__}") + + prompts = [args.prompt] + + class _SP: + max_tokens = args.max_new_tokens + temperature = 0.0 + + # First call includes prefill + any lazy engine bring-up; measure separately. + t_first0 = time.perf_counter() + outputs = model.fast_generate(prompts, sampling_params = _SP(), use_tqdm = False) + t_first = time.perf_counter() - t_first0 + out = outputs[0] + n_tok = len(out.outputs[0].token_ids) + + # Warm steady-state: run again and measure. + t_warm0 = time.perf_counter() + outputs2 = model.fast_generate(prompts, sampling_params = _SP(), use_tqdm = False) + t_warm = time.perf_counter() - t_warm0 + n_tok_warm = len(outputs2[0].outputs[0].token_ids) + + peak_after_gen = torch.cuda.max_memory_reserved() / 1024**3 + print( + f"[smoke] first call: generated {n_tok} tokens in {t_first:.2f}s " + f"({n_tok / t_first:.1f} tok/s)" + ) + print( + f"[smoke] warm call: generated {n_tok_warm} tokens in {t_warm:.2f}s " + f"({n_tok_warm / t_warm:.1f} tok/s)" + ) + print(f"[smoke] peak VRAM after gen: {peak_after_gen:.2f} GB") + print(f"[smoke] prompt: {args.prompt!r}") + print(f"[smoke] completion: {out.outputs[0].text!r}") + + precision = "4bit" if args.load_in_4bit else args.dtype + out_dir = Path(args.out_dir) + out_dir.mkdir(parents = True, exist_ok = True) + summary = { + "phase": "smoke_A", + "model": args.model, + "precision": precision, + "dtype": str(dtype), + "max_seq_length": args.max_seq_length, + "max_new_tokens": args.max_new_tokens, + "with_lora": args.with_lora, + "t_load_s": round(t_load, 2), + "peak_vram_after_load_gb": round(peak_after_load, 2), + "peak_vram_after_gen_gb": round(peak_after_gen, 2), + "first_call_s": round(t_first, 2), + "first_call_tok_s": round(n_tok / t_first, 1), + "warm_call_s": round(t_warm, 2), + "warm_call_tok_s": round(n_tok_warm / t_warm, 1), + "arch": arch, + "impl": impl, + "prompt": args.prompt, + "completion": out.outputs[0].text, + } + with open(out_dir / f"smoke_A_{precision}.json", "w") as f: + json.dump(summary, f, indent = 2) + print(f"[smoke] wrote {out_dir / f'smoke_A_{precision}.json'}") + + +if __name__ == "__main__": + main() diff --git a/unsloth/inference/flex_engine.py b/unsloth/inference/flex_engine.py index 935ba8e31c..0559703c95 100644 --- a/unsloth/inference/flex_engine.py +++ b/unsloth/inference/flex_engine.py @@ -765,14 +765,32 @@ class FlexEngine: # Materialise the pristine source the first time we see a LoRA. if self._pristine_base is None: - # The inference model has already been flex-patched; its - # linear weights (what LoRA merges into) are still pristine, - # so we can clone it and just not call flex attention on the - # pristine copy. - with weight_pool(self._cumem_allocator): - self._pristine_base = copy.deepcopy(self._inference_model) - self._pristine_base.eval() - self._impl.base_model = self._pristine_base + if self.arch == "qwen3_moe": + # For Qwen3 MoE, LoRA lives on the stacked-expert + # ParamWrapper and never merges in-place into the expert + # tensors during training (see + # unsloth_zoo.temporary_patches.moe_utils + # _patched_param_wrapper_forward). That means the + # training model's expert weights ARE the pristine + # source — no third 30-60 GB deep-copy is needed. Point + # the LoRA-refresh helper at the training model's base + # directly. This keeps the 30B MoE model at 2x residency + # instead of 3x. + try: + pristine = training_peft_model.get_base_model() + except AttributeError: + pristine = training_peft_model + self._pristine_base = pristine + self._impl.base_model = pristine + else: + # Dense path: the inference model has already been + # flex-patched; its linear weights (what LoRA merges + # into) are still pristine, so we clone it and just do + # not call flex attention on the pristine copy. + with weight_pool(self._cumem_allocator): + self._pristine_base = copy.deepcopy(self._inference_model) + self._pristine_base.eval() + self._impl.base_model = self._pristine_base if self._inference_peft is None: try: diff --git a/unsloth/inference/flex_moe.py b/unsloth/inference/flex_moe.py index 23c27f05f1..6be8d69329 100644 --- a/unsloth/inference/flex_moe.py +++ b/unsloth/inference/flex_moe.py @@ -85,25 +85,30 @@ def call_moe_model_with_flex_kwargs(model, input_ids, position_ids, flex_kwargs) _sin = _sin[position_ids] position_embeddings = (_cos, _sin) hidden_states = inputs_embeds + # RMSNorm + bnb-4bit Linear compute can promote activations to fp32 + # along the Qwen3 MoE path even under autocast. Lock activations to + # the embed dtype so paged-KV writes (which index_put_ into a + # pre-allocated bf16 cache) see a matching dtype. + compute_dtype = inputs_embeds.dtype for layer in base.layers: # Attention block — identical to dense Qwen3 / Llama. residual = hidden_states - hidden_states = layer.input_layernorm(hidden_states) + hidden_states = layer.input_layernorm(hidden_states).to(compute_dtype) hidden_states, _ = layer.self_attn( hidden_states, position_embeddings = position_embeddings, **flex_kwargs, ) - hidden_states = residual + hidden_states + hidden_states = residual + hidden_states.to(compute_dtype) # MoE MLP. residual = hidden_states - hidden_states = layer.post_attention_layernorm(hidden_states) + hidden_states = layer.post_attention_layernorm(hidden_states).to(compute_dtype) mlp_out = layer.mlp(hidden_states) if isinstance(mlp_out, tuple): hidden_states = mlp_out[0] else: hidden_states = mlp_out - hidden_states = residual + hidden_states + hidden_states = residual + hidden_states.to(compute_dtype) hidden_states = base.norm(hidden_states) return hidden_states @@ -136,6 +141,41 @@ class FlexMoEInference: peft_model = None, cumem_allocator = None, ): + # FastQwen3MoeModel.pre_patch (unsloth/models/qwen3_moe.py) installs + # a legacy Qwen3MoeSparseMoeBlock_fast_forward that expects + # ``self.gate_proj``; transformers 5.x Qwen3MoE uses + # ``self.gate`` / ``self.experts`` instead, so that forward is dead + # code on this env. Unsloth-zoo's ``patch_qwen3_moe`` re-patches it + # to the correct ``sparse_moe_block_forward``, but Unsloth's + # pre_patch can run later and silently clobber it (patch_function + # bails via can_safely_patch on a second pass). Force-restore the + # stock HF forward here so the flex walker sees a working MLP. + try: + import transformers.models.qwen3_moe.modeling_qwen3_moe as _hf_mod + _BlockCls = _hf_mod.Qwen3MoeSparseMoeBlock + cur_forward = getattr(_BlockCls, "forward", None) + cur_name = getattr(cur_forward, "__name__", "") + if "fast_forward" in cur_name or cur_name == "Qwen3MoeSparseMoeBlock_fast_forward": + # Prefer unsloth_zoo's patched version if present; + # fall back to the stock HF forward otherwise. + unique = getattr(_BlockCls, "_original_forward_Qwen3MoeSparseMoeBlock", None) or getattr(_BlockCls, "_Qwen3MoeSparseMoeBlock_original_forward", None) + if unique is not None: + _BlockCls.forward = unique + else: + # Re-run unsloth_zoo patch to install sparse_moe_block_forward. + from unsloth_zoo.temporary_patches.qwen3_moe import patch_qwen3_moe + patch_qwen3_moe() + # If patch_function still skipped due to can_safely_patch, + # fall back to stock HF as a last resort. + cur_forward_after = getattr(_BlockCls, "forward", None) + cur_name_after = getattr(cur_forward_after, "__name__", "") + if "fast_forward" in cur_name_after: + # Lazy-load pristine forward by reloading the module. + import importlib + _fresh_mod = importlib.reload(_hf_mod) + _BlockCls.forward = _fresh_mod.Qwen3MoeSparseMoeBlock.forward + except Exception: + pass assert max_seq_length % page_size == 0 # Startup sanity checks. If any of these fail the architecture # isn't a Qwen3-MoE variant we know how to drive. diff --git a/unsloth/inference/flex_qwen3_llama.py b/unsloth/inference/flex_qwen3_llama.py index 4d38cd2cbe..a13caf449a 100644 --- a/unsloth/inference/flex_qwen3_llama.py +++ b/unsloth/inference/flex_qwen3_llama.py @@ -161,7 +161,15 @@ def make_flex_attention_forward(page_table: PageTable): # Write to paged KV cache. For prefill, assign_prefill_no_paging # writes into [1, H, MAX_S, D]; for decode, assign() writes into the # B decode slots. + # Match the pre-allocated KV cache dtype; bnb-4bit Linear compute + # can produce fp32 k/v even under autocast, and the paged-cache + # index_put_ refuses mixed dtypes. if self._paged_cache is not None and flex_input_pos is not None: + cache_dtype = self._paged_cache.k_cache.dtype + if k.dtype != cache_dtype: + k = k.to(cache_dtype) + if v.dtype != cache_dtype: + v = v.to(cache_dtype) k, v = self._paged_cache.update(flex_input_pos, k, v, flex_batch_idx) # Flex attention. The block mask routes each query to the correct diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 93a1592a5e..57733aab40 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -3408,7 +3408,7 @@ class FastLlamaModel: apply_lora_mlp = apply_lora_mlp_swiglu elif model_type == "falcon_h1": apply_lora_mlp = apply_lora_mlp_swiglu - elif model_type == "qwen3moe": + elif model_type == "qwen3_moe": apply_lora_mlp = apply_lora_mlp_swiglu else: raise NotImplementedError(f"Unsloth: {model_type} is not yet implemented!") @@ -3482,6 +3482,16 @@ class FastLlamaModel: # MLP patching mlp_module = layer.mlp + # Qwen3 MoE uses Qwen3MoeSparseMoeBlock which holds + # stacked expert tensors on .experts; the dense + # gate/up/down fusion does not apply. MoE LoRA is + # wired through unsloth_zoo/moe_utils instead. + if not ( + hasattr(mlp_module, "gate_proj") + and hasattr(mlp_module, "up_proj") + and hasattr(mlp_module, "down_proj") + ): + continue gate_proj = mlp_module.gate_proj up_proj = mlp_module.up_proj down_proj = mlp_module.down_proj