diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml index 80cce7793b..fc864d1736 100644 --- a/.github/workflows/stale.yml +++ b/.github/workflows/stale.yml @@ -11,7 +11,7 @@ jobs: issues: write steps: - - uses: actions/stale@v9 + - uses: actions/stale@v10 with: # The message to post on stale issues. # This message will ping the issue author. diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index bd37ece943..ec6fb860dc 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,6 +1,6 @@ repos: - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.14.13 + rev: v0.14.14 hooks: - id: ruff args: diff --git a/README.md b/README.md index ff8dcdeef6..2b9f1cabb0 100644 --- a/README.md +++ b/README.md @@ -23,18 +23,18 @@ Notebooks are beginner friendly. Read our [guide](https://unsloth.ai/docs/get-st | Model | Free Notebooks | Performance | Memory use | |-----------|---------|--------|----------| | **gpt-oss (20B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/gpt-oss-(20B)-Fine-tuning.ipynb) | 1.5x faster | 70% less | -| **Mistral Ministral 3 (3B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Ministral_3_VL_(3B)_Vision.ipynb) | 1.5x faster | 60% less | | **gpt-oss (20B): GRPO** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/gpt-oss-(20B)-GRPO.ipynb) | 2x faster | 80% less | | **Qwen3: Advanced GRPO** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_(4B)-GRPO.ipynb) | 2x faster | 50% less | | **Qwen3-VL (8B): GSPO** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_VL_(8B)-Vision-GRPO.ipynb) | 1.5x faster | 80% less | -| **Gemma 3 (270M)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Gemma3_(270M).ipynb) | 1.7x faster | 60% less | -| **Gemma 3n (4B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Gemma3N_(4B)-Conversational.ipynb) | 1.5x faster | 50% less | -| **DeepSeek-OCR (3B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Deepseek_OCR_(3B).ipynb) | 1.5x faster | 30% less | +| **Gemma 3 (4B) Vision** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Gemma3_(4B)-Vision.ipynb) | 1.7x faster | 60% less | +| **Gemma 3n (e4B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Gemma3N_(4B)-Conversational.ipynb) | 1.5x faster | 50% less | +| **embeddinggemma (300M)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/EmbeddingGemma_(300M).ipynb) | 2x faster | 20% less | +| **Mistral Ministral 3 (3B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Ministral_3_VL_(3B)_Vision.ipynb) | 1.5x faster | 60% less | | **Llama 3.1 (8B) Alpaca** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3.1_(8B)-Alpaca.ipynb) | 2x faster | 70% less | | **Llama 3.2 Conversational** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3.2_(1B_and_3B)-Conversational.ipynb) | 2x faster | 70% less | | **Orpheus-TTS (3B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Orpheus_(3B)-TTS.ipynb) | 1.5x faster | 50% less | -- See all our notebooks for: [Kaggle](https://github.com/unslothai/notebooks?tab=readme-ov-file#-kaggle-notebooks), [GRPO](https://unsloth.ai/docs/get-started/unsloth-notebooks#grpo-reasoning-rl-notebooks), [TTS](https://unsloth.ai/docs/get-started/unsloth-notebooks#text-to-speech-tts-notebooks) & [Vision](https://unsloth.ai/docs/get-started/unsloth-notebooks#vision-multimodal-notebooks) +- See all our notebooks for: [Kaggle](https://github.com/unslothai/notebooks?tab=readme-ov-file#-kaggle-notebooks), [GRPO](https://unsloth.ai/docs/get-started/unsloth-notebooks#grpo-reasoning-rl-notebooks), [TTS](https://unsloth.ai/docs/get-started/unsloth-notebooks#text-to-speech-tts-notebooks), [embedding](https://unsloth.ai/docs/new/embedding-finetuning) & [Vision](https://unsloth.ai/docs/get-started/unsloth-notebooks#vision-multimodal-notebooks) - See [all our models](https://unsloth.ai/docs/get-started/unsloth-model-catalog) and [all our notebooks](https://unsloth.ai/docs/get-started/unsloth-notebooks) - See detailed documentation for Unsloth [here](https://unsloth.ai/docs) @@ -53,22 +53,22 @@ Use our official [Unsloth Docker image](https://hub.docker.com/r/unsloth/unsloth For RTX 50x, B200, 6000 GPUs: `pip install unsloth`. Read our [Blackwell Guide](https://unsloth.ai/docs/basics/fine-tuning-llms-with-blackwell-rtx-50-series-and-unsloth) and [DGX Spark Guide](https://unsloth.ai/docs/basics/fine-tuning-llms-with-nvidia-dgx-spark-and-unsloth) for more details. ## 🦥 Unsloth News -- New 7x longer context reinforcement learning vs. all other setups, via our new batching algorithms. [Blog](https://unsloth.ai/docs/new/grpo-long-context) +- **Embedding models**: Unsloth now supports ~1.8-3.3x faster embedding fine-tuning. [Blog](https://unsloth.ai/docs/new/embedding-finetuning) • [Notebooks](https://unsloth.ai/docs/get-started/unsloth-notebooks#embedding-models) +- New **7x longer context RL** vs. all other setups, via our new batching algorithms. [Blog](https://unsloth.ai/docs/new/grpo-long-context) - New RoPE & MLP **Triton Kernels** & **Padding Free + Packing**: 3x faster training & 30% less VRAM. [Blog](https://unsloth.ai/docs/new/3x-faster-training-packing) -- **Mistral 3**: Run Ministral 3 or Devstral 2 and fine-tune with vision/RL sodoku notebooks. [Guide](https://unsloth.ai/docs/models/ministral-3) • [Notebooks](https://unsloth.ai/docs/models/ministral-3#fine-tuning-ministral-3) - **500K Context**: Training a 20B model with >500K context is now possible on an 80GB GPU. [Blog](https://unsloth.ai/docs/new/500k-context-length-fine-tuning) - **FP8 Reinforcement Learning**: You can now do FP8 GRPO on consumer GPUs. [Blog](https://unsloth.ai/docs/new/fp8-reinforcement-learning) • [Notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_8B_FP8_GRPO.ipynb) - **DeepSeek-OCR**: Fine-tune to improve language understanding by 89%. [Guide](https://unsloth.ai/docs/models/deepseek-ocr-how-to-run-and-fine-tune) • [Notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Deepseek_OCR_(3B).ipynb) - **Docker**: Use Unsloth with no setup & environment issues with our new image. [Guide](https://unsloth.ai/docs/new/how-to-fine-tune-llms-with-unsloth-and-docker) • [Docker image](https://hub.docker.com/r/unsloth/unsloth) -- **gpt-oss RL**: Introducing the fastest possible inference for gpt-oss RL! [Read blog](https://unsloth.ai/docs/models/gpt-oss-how-to-run-and-fine-tune/gpt-oss-reinforcement-learning) - **Vision RL**: You can now train VLMs with GRPO or GSPO in Unsloth! [Read guide](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/vision-reinforcement-learning-vlm-rl) -- **gpt-oss** by OpenAI: Read our [Unsloth Flex Attention](https://unsloth.ai/docs/models/gpt-oss-how-to-run-and-fine-tune/long-context-gpt-oss-training) blog and [gpt-oss Guide](https://unsloth.ai/docs/models/gpt-oss-how-to-run-and-fine-tune). 20B works on 14GB VRAM. 120B on 65GB. +- **gpt-oss** by OpenAI: Read our [RL blog](https://unsloth.ai/docs/models/gpt-oss-how-to-run-and-fine-tune/gpt-oss-reinforcement-learning), [Flex Attention](https://unsloth.ai/docs/models/gpt-oss-how-to-run-and-fine-tune/long-context-gpt-oss-training) blog and [gpt-oss Guide](https://unsloth.ai/docs/models/gpt-oss-how-to-run-and-fine-tune). 20B works on 14GB VRAM. 120B on 65GB.
Click for more news - **Quantization-Aware Training**: We collabed with Pytorch, recovering ~70% accuracy. [Read blog](https://unsloth.ai/docs/basics/quantization-aware-training-qat) - **Memory-efficient RL**: We're introducing even better RL. Our new kernels & algos allows faster RL with 50% less VRAM & 10× more context. [Read blog](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/memory-efficient-rl) +- **Mistral 3**: Run Ministral 3 or Devstral 2 and fine-tune with vision/RL sodoku notebooks. [Guide](https://unsloth.ai/docs/models/ministral-3) • [Notebooks](https://unsloth.ai/docs/models/ministral-3#fine-tuning-ministral-3) - **Gemma 3n** by Google: [Read Blog](https://unsloth.ai/docs/models/gemma-3-how-to-run-and-fine-tune/gemma-3n-how-to-run-and-fine-tune). We [uploaded GGUFs, 4-bit models](https://huggingface.co/collections/unsloth/gemma-3n-685d3874830e49e1c93f9339). - **[Text-to-Speech (TTS)](https://unsloth.ai/docs/basics/text-to-speech-tts-fine-tuning)** is now supported, including `sesame/csm-1b` and STT `openai/whisper-large-v3`. - **[Qwen3](https://unsloth.ai/docs/models/qwen3-how-to-run-and-fine-tune)** is now supported. Qwen3-30B-A3B fits on 17.5GB VRAM. @@ -99,7 +99,7 @@ For RTX 50x, B200, 6000 GPUs: `pip install unsloth`. Read our [Blackwell Guide]( ## ⭐ Key Features * Supports **full-finetuning**, pretraining, 4b-bit, 16-bit and **FP8** training -* Supports **all models** including [TTS](https://unsloth.ai/docs/basics/text-to-speech-tts-fine-tuning), multimodal, [BERT](https://unsloth.ai/docs/get-started/unsloth-notebooks#other-important-notebooks) and more! Any model that works in transformers, works in Unsloth. +* Supports **all models** including [TTS](https://unsloth.ai/docs/basics/text-to-speech-tts-fine-tuning), multimodal, [embedding](https://unsloth.ai/docs/new/embedding-finetuning) and more! Any model that works in transformers, works in Unsloth. * The most efficient library for [Reinforcement Learning (RL)](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide), using 80% less VRAM. Supports GRPO, GSPO, DrGRPO, DAPO etc. * **0% loss in accuracy** - no approximation methods - all exact. * Export and [deploy your model](https://unsloth.ai/docs/basics/inference-and-deployment) to GGUF, llama.cpp, vLLM, SGLang and Hugging Face. @@ -194,9 +194,9 @@ pip install unsloth
### Advanced Pip Installation -`⚠️Do **NOT** use this if you have Conda.` Pip is a bit more complex since there are dependency issues. The pip command is different for `torch 2.2,2.3,2.4,2.5,2.6,2.7,2.8,2.9` and CUDA versions. +`⚠️Do **NOT** use this if you have Conda.` Pip is a bit more complex since there are dependency issues. The pip command is different for `torch 2.2,2.3,2.4,2.5,2.6,2.7,2.8,2.9,2.10` and CUDA versions. -For other torch versions, we support `torch211`, `torch212`, `torch220`, `torch230`, `torch240`, `torch250`, `torch260`, `torch270`, `torch280`, `torch290` and for CUDA versions, we support `cu118` and `cu121` and `cu124`. For Ampere devices (A100, H100, RTX3090) and above, use `cu118-ampere` or `cu121-ampere` or `cu124-ampere`. +For other torch versions, we support `torch211`, `torch212`, `torch220`, `torch230`, `torch240`, `torch250`, `torch260`, `torch270`, `torch280`, `torch290`, `torch2100` and for CUDA versions, we support `cu118` and `cu121` and `cu124`. For Ampere devices (A100, H100, RTX3090) and above, use `cu118-ampere` or `cu121-ampere` or `cu124-ampere`. Note: torch 2.10 only supports CUDA 12.6, 12.8, and 13.0. For example, if you have `torch 2.4` and `CUDA 12.1`, use: ```bash @@ -210,6 +210,12 @@ pip install --upgrade pip pip install "unsloth[cu130-torch290] @ git+https://github.com/unslothai/unsloth.git" ``` +Another example, if you have `torch 2.10` and `CUDA 12.6`, use: +```bash +pip install --upgrade pip +pip install "unsloth[cu126-torch2100] @ git+https://github.com/unslothai/unsloth.git" +``` + And other examples: ```bash pip install "unsloth[cu121-ampere-torch240] @ git+https://github.com/unslothai/unsloth.git" @@ -254,8 +260,10 @@ elif v < V('2.8.0'): x = 'cu{}{}-torch271' elif v < V('2.8.9'): x = 'cu{}{}-torch280' elif v < V('2.9.1'): x = 'cu{}{}-torch290' elif v < V('2.9.2'): x = 'cu{}{}-torch291' +elif v < V('2.10.1'): x = 'cu{}{}-torch2100' else: raise RuntimeError(f"Torch = {v} too new!") if v > V('2.6.9') and cuda not in ("11.8", "12.6", "12.8", "13.0"): raise RuntimeError(f"CUDA = {cuda} not supported!") +if v >= V('2.10.0') and cuda not in ("12.6", "12.8", "13.0"): raise RuntimeError(f"Torch 2.10 requires CUDA 12.6, 12.8, or 13.0! Got CUDA = {cuda}") x = x.format(cuda.replace(".", ""), "-ampere" if False else "") # is_ampere is broken due to flash-attn print(f'pip install --upgrade pip && pip install --no-deps git+https://github.com/unslothai/unsloth-zoo.git && pip install "unsloth[{x}] @ git+https://github.com/unslothai/unsloth.git" --no-build-isolation') ``` diff --git a/unsloth/__init__.py b/unsloth/__init__.py index d3093cf4c0..fad37a786d 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -79,7 +79,7 @@ from importlib.metadata import PackageNotFoundError # Check for unsloth_zoo try: unsloth_zoo_version = importlib_version("unsloth_zoo") - if Version(unsloth_zoo_version) < Version("2026.1.2"): + if Version(unsloth_zoo_version) < Version("2026.2.1"): print( "Unsloth: Please update Unsloth and Unsloth-Zoo to the latest version!\n" "Do this via `pip install --upgrade --force-reinstall --no-cache-dir --no-deps unsloth unsloth_zoo`" @@ -125,43 +125,56 @@ from unsloth_zoo.device_type import ( from .import_fixes import ( fix_xformers_performance_issue, fix_vllm_aimv2_issue, + check_vllm_torch_sm100_compatibility, fix_vllm_guided_decoding_params, fix_vllm_pdl_blackwell, + fix_rocm_triton_key_error, ignore_logger_messages, patch_ipykernel_hf_xet, patch_trackio, patch_datasets, patch_enable_input_require_grads, fix_openenv_no_vllm, + patch_openspiel_env_async, fix_executorch, patch_vllm_for_notebooks, + patch_torchcodec_audio_decoder, ) fix_xformers_performance_issue() fix_vllm_aimv2_issue() +# Check vLLM + torch < 2.9.0 + SM100 compatibility BEFORE importing vLLM +check_vllm_torch_sm100_compatibility() fix_vllm_guided_decoding_params() fix_vllm_pdl_blackwell() +fix_rocm_triton_key_error() ignore_logger_messages() patch_ipykernel_hf_xet() patch_trackio() patch_datasets() patch_enable_input_require_grads() fix_openenv_no_vllm() +patch_openspiel_env_async() fix_executorch() patch_vllm_for_notebooks() +patch_torchcodec_audio_decoder() del fix_xformers_performance_issue del fix_vllm_aimv2_issue +del check_vllm_torch_sm100_compatibility del fix_vllm_guided_decoding_params del fix_vllm_pdl_blackwell +del fix_rocm_triton_key_error del ignore_logger_messages del patch_ipykernel_hf_xet del patch_trackio del patch_datasets del patch_enable_input_require_grads del fix_openenv_no_vllm +del patch_openspiel_env_async del fix_executorch del patch_vllm_for_notebooks +del patch_torchcodec_audio_decoder # Torch 2.4 has including_emulation if DEVICE_TYPE == "cuda": diff --git a/unsloth/_auto_install.py b/unsloth/_auto_install.py index ed9b918412..f6deefeb33 100644 --- a/unsloth/_auto_install.py +++ b/unsloth/_auto_install.py @@ -35,7 +35,9 @@ elif v < V('2.8.0'): x = 'cu{}{}-torch271' elif v < V('2.8.9'): x = 'cu{}{}-torch280' elif v < V('2.9.1'): x = 'cu{}{}-torch290' elif v < V('2.9.2'): x = 'cu{}{}-torch291' +elif v < V('2.10.1'): x = 'cu{}{}-torch2100' else: raise RuntimeError(f"Torch = {v} too new!") if v > V('2.6.9') and cuda not in ("11.8", "12.6", "12.8", "13.0"): raise RuntimeError(f"CUDA = {cuda} not supported!") +if v >= V('2.10.0') and cuda not in ("12.6", "12.8", "13.0"): raise RuntimeError(f"Torch 2.10 requires CUDA 12.6, 12.8, or 13.0! Got CUDA = {cuda}") x = x.format(cuda.replace(".", ""), "-ampere" if False else "") # is_ampere is broken due to flash-attn print(f'pip install --upgrade pip && pip install --no-deps git+https://github.com/unslothai/unsloth-zoo.git && pip install "unsloth[{x}] @ git+https://github.com/unslothai/unsloth.git" --no-build-isolation') \ No newline at end of file diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py index 4f88808c2a..4fbab6a94c 100644 --- a/unsloth/import_fixes.py +++ b/unsloth/import_fixes.py @@ -123,6 +123,47 @@ if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1": warnings.filterwarnings("ignore", message = "`int4_weight_only` is deprecated") warnings.filterwarnings("ignore", message = "`int8_weight_only` is deprecated") + # TorchAO deprecated import paths (https://github.com/pytorch/ao/issues/2752) + warnings.filterwarnings( + "ignore", + message = r"Importing.*from torchao\.dtypes.*is deprecated", + category = DeprecationWarning, + ) + warnings.filterwarnings( + "ignore", + message = r"Importing BlockSparseLayout from torchao\.dtypes is deprecated", + category = DeprecationWarning, + ) + + # SWIG builtin type warnings (from bitsandbytes/triton SWIG bindings) + warnings.filterwarnings( + "ignore", + message = r"builtin type Swig.*has no __module__ attribute", + category = DeprecationWarning, + ) + + # Triton autotuner deprecation (https://github.com/triton-lang/triton/pull/4496) + warnings.filterwarnings( + "ignore", + message = r"warmup, rep, and use_cuda_graph parameters are deprecated", + category = DeprecationWarning, + ) + + # Python 3.12+ multiprocessing fork warning in multi-threaded processes + warnings.filterwarnings( + "ignore", + message = r".*multi-threaded.*use of fork\(\) may lead to deadlocks", + category = DeprecationWarning, + ) + + # Resource warnings from internal socket/file operations + warnings.filterwarnings( + "ignore", message = r"unclosed.*socket", category = ResourceWarning + ) + warnings.filterwarnings( + "ignore", message = r"unclosed file.*dev/null", category = ResourceWarning + ) + # Fix up AttributeError: 'MessageFactory' object has no attribute 'GetPrototype' # MUST do this at the start primarily due to tensorflow causing issues @@ -503,46 +544,138 @@ def patch_enable_input_require_grads(): ) +def _is_custom_torch_build(raw_version_str): + """Check if a raw version string indicates a custom or source build. + Must operate on the raw string from importlib_version(), not the parsed + Version object, since our custom Version() strips local identifiers. + + Standard PyTorch releases use: +cu124, +rocm6.3, +cpu, +xpu + Source/custom builds use: +gitXXXXXXX, +HEXHASH, or other suffixes. + """ + if "+" not in raw_version_str: + return False + local = raw_version_str.split("+", 1)[1] + if not local: + return False + # Use fullmatch so the entire local identifier must match, not just a prefix. + # cu/rocm require a trailing digit (e.g. cu124, rocm6.3). cpu/xpu are exact. + # Case-insensitive since some builds may use uppercase. + return not re.fullmatch(r"cu\d[\d.]*|rocm\d[\d.]*|cpu|xpu", local, re.IGNORECASE) + + +def _infer_required_torchvision(torch_major, torch_minor): + """Infer the minimum required torchvision minor version from torch version. + + The torch -> torchvision minor version mapping follows a consistent formula: + torch 1.x -> torchvision 0.(x + 1) (verified: torch 1.7 through 1.13) + torch 2.x -> torchvision 0.(x + 15) (verified: torch 2.0 through 2.9) + + Returns (tv_major, tv_minor) or None if the major version is unrecognized. + """ + if torch_major == 1 and torch_minor >= 7: + return (0, torch_minor + 1) + if torch_major == 2: + return (0, torch_minor + 15) + return None + + def torchvision_compatibility_check(): + # Allow skipping via environment variable for custom environments + if os.environ.get("UNSLOTH_SKIP_TORCHVISION_CHECK", "0").lower() in ("1", "true"): + return + if importlib.util.find_spec("torch") is None: raise ImportError("Unsloth: torch not found. Please install torch first.") if importlib.util.find_spec("torchvision") is None: return - torch_version = importlib_version("torch") - torchvision_version = importlib_version("torchvision") - # Torch version -> minimum required torchvision version - # See https://pytorch.org/get-started/previous-versions/ - TORCH_TORCHVISION_COMPAT = [ - ("2.9.0", "0.24.0"), - ("2.8.0", "0.23.0"), - ("2.7.0", "0.22.0"), - ("2.6.0", "0.21.0"), - ("2.5.0", "0.20.0"), - ("2.4.0", "0.19.0"), - ] - - required_torchvision = None - for min_torch, min_torchvision in TORCH_TORCHVISION_COMPAT: - if Version(torch_version) >= Version(min_torch): - required_torchvision = min_torchvision - break - - if required_torchvision is None: - # Torch version not in compatibility table, skip check + try: + torch_version_raw = importlib_version("torch") + torchvision_version_raw = importlib_version("torchvision") + except Exception: return - if Version(torchvision_version) < Version(required_torchvision): - raise ImportError( - f"Unsloth: torch=={torch_version} requires torchvision>={required_torchvision}, " - f"but found torchvision=={torchvision_version}. " - f"Please refer to https://pytorch.org/get-started/previous-versions/ for more information." - ) + try: + torch_v = Version(torch_version_raw) + tv_v = Version(torchvision_version_raw) + except Exception: + return - logger.info( - f"Unsloth: torch=={torch_version} and torchvision=={torchvision_version} are compatible." + # Known compatibility table (ground truth, takes precedence over formula). + # See https://pytorch.org/get-started/previous-versions/ + TORCH_TORCHVISION_COMPAT = { + (2, 9): (0, 24), + (2, 8): (0, 23), + (2, 7): (0, 22), + (2, 6): (0, 21), + (2, 5): (0, 20), + (2, 4): (0, 19), + } + + # Extract major.minor from the parsed version + torch_release = torch_v.release + if len(torch_release) < 2: + return + torch_major, torch_minor = torch_release[0], torch_release[1] + + # Try known table first, then fall back to formula for forward compatibility + required = TORCH_TORCHVISION_COMPAT.get((torch_major, torch_minor)) + is_in_known_table = required is not None + + if required is None: + required = _infer_required_torchvision(torch_major, torch_minor) + + if required is None: + return + + required_tv_str = f"{required[0]}.{required[1]}.0" + + if tv_v >= Version(required_tv_str): + logger.info( + f"Unsloth: torch=={torch_version_raw} and " + f"torchvision=={torchvision_version_raw} are compatible." + ) + return + + # Version mismatch detected + message = ( + f"Unsloth: torch=={torch_version_raw} requires " + f"torchvision>={required_tv_str}, " + f"but found torchvision=={torchvision_version_raw}. " + f"Please refer to https://pytorch.org/get-started/previous-versions/ " + f"for more information." ) + is_custom = _is_custom_torch_build(torch_version_raw) or _is_custom_torch_build( + torchvision_version_raw + ) + + # Detect nightly/dev/alpha/beta/rc builds from the raw version string. + # These often have version mismatches that are expected. + _pre_tags = (".dev", "a0", "b0", "rc", "alpha", "beta", "nightly") + is_prerelease = any(t in torch_version_raw for t in _pre_tags) or any( + t in torchvision_version_raw for t in _pre_tags + ) + + # Downgrade to warning for custom/source/pre-release builds or formula-predicted + if is_custom or is_prerelease or not is_in_known_table: + reason = ( + "custom/source build" + if is_custom + else "pre-release build" + if is_prerelease + else "newer torch version" + ) + logger.warning( + f"{message}\n" + f"Detected a {reason}. " + f"Continuing with a warning. " + f"Set UNSLOTH_SKIP_TORCHVISION_CHECK=1 to silence this." + ) + return + + raise ImportError(message) + # Fix TRL OpenEnv 0.26 NameError: name 'SamplingParams' is not defined def fix_openenv_no_vllm(): @@ -666,6 +799,105 @@ def fix_huggingface_hub(): ) +def fix_rocm_triton_key_error(): + """ + ROCm + torch.compile can fail if Triton lacks `triton_key`. + Disable Inductor/compile only on ROCm when that symbol is missing. + """ + try: + import torch + except (ImportError, ModuleNotFoundError): + return + + if not getattr(torch.version, "hip", None): + return + + try: + import triton + except (ImportError, ModuleNotFoundError): + return + + try: + from triton.runtime import triton_key # noqa: F401 + + return + except ImportError: + pass + + os.environ.setdefault("TORCHINDUCTOR_DISABLE", "1") + os.environ.setdefault("TORCH_COMPILE_DISABLE", "1") + logger.info( + "Unsloth: ROCm detected and Triton lacks triton_key; " + "disabling torch.compile/Inductor to avoid backend crash." + ) + + +def check_vllm_torch_sm100_compatibility(): + """ + Check for incompatible vLLM + torch < 2.9.0 + SM100 (Blackwell) combination. + + vLLM's distributed module (device_communicators) crashes with std::bad_alloc + when imported on SM100 GPUs (B200/B100) with torch < 2.9.0. This is due to + C++ code in vLLM's NCCL/distributed layer being incompatible with older + torch versions on the newer Blackwell architecture. + + This check runs early (before vLLM import) to provide a helpful error message + instead of a cryptic std::bad_alloc crash. + """ + # Check if vLLM is installed (without importing it) + if importlib.util.find_spec("vllm") is None: + return + + # Check torch version + try: + torch_version = Version(importlib_version("torch")) + if torch_version >= Version("2.9.0"): + return # torch >= 2.9.0 is compatible + except Exception: + return # Can't determine torch version, skip check + + # Check if any CUDA GPU is SM100 (Blackwell) + try: + import torch + + if not torch.cuda.is_available(): + return + + has_sm100 = False + sm100_gpu_name = None + for i in range(torch.cuda.device_count()): + major, minor = torch.cuda.get_device_capability(i) + if major == 10: + has_sm100 = True + sm100_gpu_name = torch.cuda.get_device_name(i) + break + + if not has_sm100: + return + except Exception: + return + + # Get vLLM version for the error message + try: + vllm_version = importlib_version("vllm") + except Exception: + vllm_version = "unknown" + + # Incompatible combination detected - raise helpful error + raise RuntimeError( + f"Unsloth: Incompatible configuration detected.\n\n" + f" GPU: {sm100_gpu_name} (SM100 / Blackwell architecture)\n" + f" torch version: {torch_version}\n" + f" vLLM version: {vllm_version}\n\n" + f"vLLM's distributed module crashes with std::bad_alloc on SM100 GPUs " + f"(B200/B100/Blackwell) when using torch < 2.9.0.\n\n" + f"To fix this, please upgrade torch:\n" + f" pip install --upgrade torch>=2.9.0\n\n" + f"Alternatively, if you don't need vLLM:\n" + f" pip uninstall vllm" + ) + + def fix_vllm_pdl_blackwell(): """ Fix vLLM PDL (Programmatic Dependent Launch) bug on Blackwell GPUs (SM100). @@ -779,3 +1011,45 @@ def fix_vllm_pdl_blackwell(): else: # Just set the env var - vLLM might be an older version without supports_pdl logger.info(f"Unsloth: Set TRITON_DISABLE_PDL=1 for SM100 ({sm100_gpu_name})") + + +def patch_openspiel_env_async(): + """Apply nest_asyncio for OpenEnv EnvClient async compatibility. + + OpenEnv's EnvClient uses async methods (reset/step). In Jupyter notebooks + these work via top-level await, but converted scripts need + asyncio.get_event_loop().run_until_complete() wrappers. Applying nest_asyncio + ensures nested event loop calls work in all contexts without replacing the + original async methods (which would break scripts that already have their own + sync wrappers). + """ + try: + import inspect + from openenv.core.env_client import EnvClient + + if not inspect.iscoroutinefunction(EnvClient.reset): + return # Already sync, nothing to do + + try: + import nest_asyncio + + nest_asyncio.apply() + logger.info( + "Unsloth: Applied nest_asyncio for OpenEnv EnvClient async compatibility" + ) + except ImportError: + logger.info( + "Unsloth: nest_asyncio not installed, OpenEnv async methods may need manual wrapping" + ) + except (ImportError, AttributeError): + pass # openenv not installed + + +def patch_torchcodec_audio_decoder(): + """Call unsloth_zoo's AudioDecoder patch.""" + try: + from unsloth_zoo.dataset_utils import patch_torchcodec_audio_decoder as _patch + + _patch() + except (ImportError, AttributeError): + pass diff --git a/unsloth/kernels/moe/autotune_cache.py b/unsloth/kernels/moe/autotune_cache.py new file mode 100644 index 0000000000..f23d9688ea --- /dev/null +++ b/unsloth/kernels/moe/autotune_cache.py @@ -0,0 +1,500 @@ +# Unsloth +# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as published +# by the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . + +""" +Auto-tuning cache system for MoE kernels to ensure tuning runs only once at training start. +""" + +import hashlib +import json +import logging +import os +import time +from typing import Dict, List, Optional, Tuple, Any +import torch +import triton + +logger = logging.getLogger(__name__) + +# Global cache for kernel configurations +_kernel_config_cache: Dict[str, Any] = {} +_autotune_completed: Dict[str, bool] = {} + + +def _get_cache_key( + num_experts: int, + hidden_dim: int, + intermediate_dim: int, + top_k: int, + dtype: torch.dtype, + device_capability: Tuple[int, int], + seq_len: int = 8192, # Default sequence length for tuning +) -> str: + """Generate a unique cache key based on model configuration.""" + key_data = { + "num_experts": num_experts, + "hidden_dim": hidden_dim, + "intermediate_dim": intermediate_dim, + "top_k": top_k, + "dtype": str(dtype), + "device_capability": device_capability, + "seq_len": seq_len, + } + key_str = json.dumps(key_data, sort_keys = True) + return hashlib.md5(key_str.encode()).hexdigest() + + +def _get_cache_file_path(cache_key: str) -> str: + """Get the file path for the cache file.""" + cache_dir = os.path.expanduser("~/.cache/unsloth/moe_autotune") + os.makedirs(cache_dir, exist_ok = True) + return os.path.join(cache_dir, f"{cache_key}.json") + + +def load_cached_config(cache_key: str) -> Optional[Dict[str, Any]]: + """Load cached kernel configuration from disk.""" + cache_file = _get_cache_file_path(cache_key) + if not os.path.exists(cache_file): + return None + + try: + with open(cache_file, "r") as f: + cached_data = json.load(f) + + # Verify cache is still valid (same device, etc.) + current_device_capability = torch.cuda.get_device_capability() + if cached_data.get("device_capability") != current_device_capability: + logger.info("Device capability changed, invalidating cache") + os.remove(cache_file) + return None + + logger.info(f"Loaded cached MoE kernel config: {cache_key}") + return cached_data + except Exception as e: + logger.warning(f"Failed to load cache file {cache_file}: {e}") + try: + os.remove(cache_file) + except: + pass + return None + + +def save_cached_config( + cache_key: str, + config_fwd: Any, + config_bwd_dx: Any, + config_bwd_dw: Any, + metadata: Dict[str, Any] = None, +) -> None: + """Save kernel configuration to disk cache.""" + cache_file = _get_cache_file_path(cache_key) + + cache_data = { + "timestamp": time.time(), + "device_capability": torch.cuda.get_device_capability(), + "config_fwd": config_fwd.__dict__ + if hasattr(config_fwd, "__dict__") + else str(config_fwd), + "config_bwd_dx": config_bwd_dx.__dict__ + if hasattr(config_bwd_dx, "__dict__") + else str(config_bwd_dx), + "config_bwd_dw": config_bwd_dw.__dict__ + if hasattr(config_bwd_dw, "__dict__") + else str(config_bwd_dw), + "metadata": metadata or {}, + } + + try: + with open(cache_file, "w") as f: + json.dump(cache_data, f, indent = 2) + logger.info(f"Saved MoE kernel config cache: {cache_key}") + except Exception as e: + logger.warning(f"Failed to save cache file {cache_file}: {e}") + + +def get_or_autotune_moe_kernels( + num_experts: int, + hidden_dim: int, + intermediate_dim: int, + top_k: int, + dtype: torch.dtype, + force_autotune: bool = False, + seq_len: int = 8192, +) -> Tuple[Any, Any, Any]: + """ + Get cached kernel configurations or run auto-tuning. + + Args: + num_experts: Number of experts in the MoE layer + hidden_dim: Hidden dimension of the model + intermediate_dim: Intermediate dimension for MoE MLP + top_k: Number of experts to route to + dtype: Data type for computation + force_autotune: Force re-running autotuning even if cache exists + seq_len: Sequence length to use for tuning benchmarks + + Returns: + Tuple of (config_fwd, config_bwd_dx, config_bwd_dw) + """ + device_capability = torch.cuda.get_device_capability() + cache_key = _get_cache_key( + num_experts, + hidden_dim, + intermediate_dim, + top_k, + dtype, + device_capability, + seq_len, + ) + + # 0. Check for environment variable override to DISABLE autotuning + if os.environ.get("UNSLOTH_MOE_DISABLE_AUTOTUNE", "0") == "1": + logger.info( + f"UNSLOTH_MOE_DISABLE_AUTOTUNE=1: Using Heuristic (Safe) MoE kernel configs for SM{device_capability[0]}{device_capability[1]}" + ) + return _get_heuristic_configs() + if not force_autotune and cache_key in _kernel_config_cache: + logger.info(f"Using in-memory cached MoE kernel configs: {cache_key}") + return _kernel_config_cache[cache_key] + + # Try to load from disk + if not force_autotune: + cached_data = load_cached_config(cache_key) + if cached_data is not None: + # Reconstruct config objects from cached data + try: + from .grouped_gemm.kernels.tuning import ( + KernelConfigForward, + KernelConfigBackward_dX, + KernelConfigBackward_dW, + ) + + config_fwd = KernelConfigForward(**cached_data["config_fwd"]) + config_bwd_dx = KernelConfigBackward_dX(**cached_data["config_bwd_dx"]) + config_bwd_dw = KernelConfigBackward_dW(**cached_data["config_bwd_dw"]) + + configs = (config_fwd, config_bwd_dx, config_bwd_dw) + _kernel_config_cache[cache_key] = configs + return configs + except Exception as e: + logger.warning(f"Failed to reconstruct cached configs: {e}") + + # Run autotuning + if cache_key in _autotune_completed and not force_autotune: + logger.info(f"Autotuning already completed for: {cache_key}") + return _kernel_config_cache[cache_key] + + logger.info(f"Running MoE kernel auto-tuning for: {cache_key}") + logger.info( + f"Configuration: {num_experts} experts, {hidden_dim} hidden, {intermediate_dim} intermediate, top_k={top_k}" + ) + + try: + configs = _run_moe_autotuning( + num_experts, hidden_dim, intermediate_dim, top_k, dtype, seq_len + ) + + # Cache the results + _kernel_config_cache[cache_key] = configs + _autotune_completed[cache_key] = True + + # Save to disk + config_fwd, config_bwd_dx, config_bwd_dw = configs + save_cached_config( + cache_key, + config_fwd, + config_bwd_dx, + config_bwd_dw, + { + "num_experts": num_experts, + "hidden_dim": hidden_dim, + "intermediate_dim": intermediate_dim, + }, + ) + + logger.info(f"MoE kernel auto-tuning completed: {cache_key}") + return configs + + except Exception as e: + logger.error(f"MoE kernel auto-tuning failed: {e}") + if "AttributeError" in str(e) and "_experimental_make_tensor_descriptor" in str( + e + ): + logger.warning( + "Unsloth: Your Triton version might be incompatible with TMA features. Falling back to default configs." + ) + logger.info("Falling back to default kernel configurations") + return _get_default_configs() + + +def _run_moe_autotuning( + num_experts: int, + hidden_dim: int, + intermediate_dim: int, + top_k: int, + dtype: torch.dtype, + seq_len: int, +) -> Tuple[Any, Any, Any]: + """Run the actual auto-tuning for MoE kernels.""" + + # Create dummy inputs for tuning + device = "cuda" + # Use a fixed, safe number of tokens for autotuning to avoid OOMs and dependency on seq_len + # 4096 is standard for finding good kernels without consuming 10GB+ VRAM + # We ignore the passed seq_len for the actual allocation to satisfy user request + num_tokens = 4096 + total_tokens = num_tokens * top_k + + # Create dummy tensors + hidden_states = torch.randn(num_tokens, hidden_dim, device = device, dtype = dtype) + + # Create dummy weights + gate_up_weights = torch.randn( + num_experts, 2 * intermediate_dim, hidden_dim, device = device, dtype = dtype + ) + down_weights = torch.randn( + num_experts, hidden_dim, intermediate_dim, device = device, dtype = dtype + ) + + # Create dummy routing data + m_sizes = torch.randint( + 1, total_tokens // num_experts + 1, (num_experts,), device = device + ) + m_sizes = m_sizes * (total_tokens // m_sizes.sum().item()) + # Adjust to ensure exact total + diff = total_tokens - m_sizes.sum().item() + if diff != 0: + m_sizes[0] += diff + + gather_indices = torch.arange(total_tokens, device = device) + torch.randperm(total_tokens, out = gather_indices) + + # Autotune forward kernel - use the interface function with autotune=True + # This properly invokes the kernel and lets triton handle the autotuning + from .grouped_gemm.interface import ( + grouped_gemm_forward, + grouped_gemm_dX, + grouped_gemm_dW, + ) + from .grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel + from .grouped_gemm.kernels.backward import ( + _autotuned_grouped_gemm_dX_kernel, + _autotuned_grouped_gemm_dW_kernel, + ) + from .grouped_gemm.kernels.tuning import ( + KernelConfigForward, + KernelConfigBackward_dX, + KernelConfigBackward_dW, + ) + + logger.info("Autotuning forward kernel (first GEMM)...") + # Run with autotune=True to trigger autotuning + _ = grouped_gemm_forward( + X = hidden_states, + W = gate_up_weights, + topk = top_k, + m_sizes = m_sizes, + gather_indices = gather_indices, + permute_x = True, + permute_y = False, + autotune = True, + ) + triton_config_fwd = _autotuned_grouped_gemm_forward_kernel.best_config + + # Convert triton.Config to KernelConfigForward + config_fwd = KernelConfigForward( + BLOCK_SIZE_M = triton_config_fwd.kwargs["BLOCK_SIZE_M"], + BLOCK_SIZE_N = triton_config_fwd.kwargs["BLOCK_SIZE_N"], + BLOCK_SIZE_K = triton_config_fwd.kwargs["BLOCK_SIZE_K"], + num_warps = triton_config_fwd.num_warps, + num_stages = triton_config_fwd.num_stages, + use_tma_load_x = triton_config_fwd.kwargs.get("USE_TMA_LOAD_X", False), + use_tma_load_w = triton_config_fwd.kwargs.get("USE_TMA_LOAD_W", False), + use_tma_store = triton_config_fwd.kwargs.get("USE_TMA_STORE", False), + ) + + # Autotune backward dX kernel + logger.info("Autotuning backward dX kernel...") + dummy_grad = torch.randn( + total_tokens, 2 * intermediate_dim, device = device, dtype = dtype + ) + _ = grouped_gemm_dX( + dY = dummy_grad, + W = gate_up_weights, + gather_indices = gather_indices, + m_sizes = m_sizes, + topk = top_k, + permute_x = True, + permute_y = False, + autotune = True, + ) + triton_config_bwd_dx = _autotuned_grouped_gemm_dX_kernel.best_config + + # Convert triton.Config to KernelConfigBackward_dX + config_bwd_dx = KernelConfigBackward_dX( + BLOCK_SIZE_M = triton_config_bwd_dx.kwargs["BLOCK_SIZE_M"], + BLOCK_SIZE_N = triton_config_bwd_dx.kwargs["BLOCK_SIZE_N"], + BLOCK_SIZE_K = triton_config_bwd_dx.kwargs["BLOCK_SIZE_K"], + num_warps = triton_config_bwd_dx.num_warps, + num_stages = triton_config_bwd_dx.num_stages, + use_tma_load_dy = triton_config_bwd_dx.kwargs.get("USE_TMA_LOAD_dY", False), + use_tma_load_w = triton_config_bwd_dx.kwargs.get("USE_TMA_LOAD_W", False), + use_tma_store = triton_config_bwd_dx.kwargs.get("USE_TMA_STORE", False), + ) + + # Autotune backward dW kernel + logger.info("Autotuning backward dW kernel...") + _ = grouped_gemm_dW( + X = hidden_states, + dY = dummy_grad, + m_sizes = m_sizes, + gather_indices = gather_indices, + topk = top_k, + permute_x = True, + permute_y = False, + autotune = True, + ) + triton_config_bwd_dw = _autotuned_grouped_gemm_dW_kernel.best_config + + # Convert triton.Config to KernelConfigBackward_dW + config_bwd_dw = KernelConfigBackward_dW( + BLOCK_SIZE_M = triton_config_bwd_dw.kwargs["BLOCK_SIZE_M"], + BLOCK_SIZE_N = triton_config_bwd_dw.kwargs["BLOCK_SIZE_N"], + BLOCK_SIZE_K = triton_config_bwd_dw.kwargs["BLOCK_SIZE_K"], + num_warps = triton_config_bwd_dw.num_warps, + num_stages = triton_config_bwd_dw.num_stages, + use_tma_load_dy = triton_config_bwd_dw.kwargs.get("USE_TMA_LOAD_dY", False), + use_tma_load_x = triton_config_bwd_dw.kwargs.get("USE_TMA_LOAD_X", False), + use_tma_store = triton_config_bwd_dw.kwargs.get("USE_TMA_STORE", False), + ) + + return config_fwd, config_bwd_dx, config_bwd_dw + + return config_fwd, config_bwd_dx, config_bwd_dw + + +def _get_heuristic_configs() -> Tuple[Any, Any, Any]: + """ + Get 'Safe Heuristic' kernel configurations. + These are verified to be safe on A100 (SM80) and provide ~9x speedup on H100/B200. + """ + from .grouped_gemm.kernels.tuning import ( + KernelConfigForward, + KernelConfigBackward_dX, + KernelConfigBackward_dW, + ) + + # Safe Forward Config: 64x128x128 (Fits A100 SMEM) + config_fwd = KernelConfigForward( + BLOCK_SIZE_M = 64, + BLOCK_SIZE_N = 128, + BLOCK_SIZE_K = 128, + num_warps = 8, + num_stages = 3, + permute_x = True, + permute_y = True, + use_tma_load_x = False, + use_tma_load_w = False, # TMA loads might need alignment checks, safer to disable for heuristic + use_tma_store = False, + ) + + # Safe Backward Configs: 64x64x256 + config_bwd_dx = KernelConfigBackward_dX( + BLOCK_SIZE_M = 64, + BLOCK_SIZE_N = 64, + BLOCK_SIZE_K = 256, + num_warps = 8, + num_stages = 4, + permute_x = True, + permute_y = True, + use_tma_load_dy = False, + use_tma_load_w = False, + use_tma_store = False, + ) + + config_bwd_dw = KernelConfigBackward_dW( + BLOCK_SIZE_M = 64, + BLOCK_SIZE_N = 64, + BLOCK_SIZE_K = 256, + num_warps = 8, + num_stages = 4, + permute_x = True, + permute_y = True, + use_tma_load_dy = False, + use_tma_load_x = False, + use_tma_store = False, + ) + + return config_fwd, config_bwd_dx, config_bwd_dw + + +def _get_default_configs() -> Tuple[Any, Any, Any]: + """Get default kernel configurations as fallback.""" + from .grouped_gemm.kernels.tuning import ( + KernelConfigForward, + KernelConfigBackward_dX, + KernelConfigBackward_dW, + ) + + logger.warning("Using default MoE kernel configurations (not optimal)") + + config_fwd = KernelConfigForward( + BLOCK_SIZE_M = 128, + BLOCK_SIZE_N = 128, + BLOCK_SIZE_K = 64, + num_warps = 8, + num_stages = 3, + use_tma_load_x = False, + use_tma_load_w = False, + use_tma_store = False, + ) + + config_bwd_dx = KernelConfigBackward_dX( + BLOCK_SIZE_M = 128, + BLOCK_SIZE_N = 128, + BLOCK_SIZE_K = 64, + num_warps = 8, + num_stages = 3, + use_tma_load_dy = False, + use_tma_load_w = False, + use_tma_store = False, + ) + + config_bwd_dw = KernelConfigBackward_dW( + BLOCK_SIZE_M = 128, + BLOCK_SIZE_N = 128, + BLOCK_SIZE_K = 64, + num_warps = 8, + num_stages = 3, + use_tma_load_dy = False, + use_tma_load_x = False, + use_tma_store = False, + ) + + return config_fwd, config_bwd_dx, config_bwd_dw + + +def clear_cache() -> None: + """Clear all cached kernel configurations.""" + global _kernel_config_cache, _autotune_completed + _kernel_config_cache.clear() + _autotune_completed.clear() + logger.info("Cleared MoE kernel cache") + + +def is_autotuning_completed(cache_key: str) -> bool: + """Check if autotuning has been completed for a given cache key.""" + return cache_key in _autotune_completed diff --git a/unsloth/kernels/moe/grouped_gemm/interface.py b/unsloth/kernels/moe/grouped_gemm/interface.py index d278d89645..572b02702b 100644 --- a/unsloth/kernels/moe/grouped_gemm/interface.py +++ b/unsloth/kernels/moe/grouped_gemm/interface.py @@ -8,17 +8,17 @@ from dataclasses import asdict import torch import triton -from grouped_gemm.kernels.backward import ( +from .kernels.backward import ( _autotuned_grouped_gemm_dW_kernel, _autotuned_grouped_gemm_dX_kernel, _grouped_gemm_dW_kernel, _grouped_gemm_dX_kernel, ) -from grouped_gemm.kernels.forward import ( +from .kernels.forward import ( _autotuned_grouped_gemm_forward_kernel, _grouped_gemm_forward_kernel, ) -from grouped_gemm.kernels.tuning import ( +from .kernels.tuning import ( KernelConfigBackward_dW, KernelConfigBackward_dX, KernelConfigForward, @@ -35,17 +35,57 @@ ch = logging.StreamHandler() ch.setFormatter(formatter) logger.addHandler(ch) -_FUSED_MUL_WARN = False -_SUPPORTS_TMA = None + +# Precompute TMA support to avoid graph breaks +# TMA requires both: +# 1. GPU capability >= 9 (Hopper+) +# 2. Triton version with TMA API (make_tensor_descriptor or _experimental_make_tensor_descriptor) +def _check_tma_support(): + import triton.language as tl + + gpu_supports_tma = torch.cuda.get_device_capability()[0] >= 9 + # Check for both old experimental and new stable API names + triton_has_tma_api = hasattr(tl, "make_tensor_descriptor") or hasattr( + tl, "_experimental_make_tensor_descriptor" + ) + return gpu_supports_tma and triton_has_tma_api + + +_SUPPORTS_TMA = _check_tma_support() + +# Check if triton.set_allocator is available (Triton 3.0+) +_HAS_SET_ALLOCATOR = hasattr(triton, "set_allocator") def supports_tma(): - global _SUPPORTS_TMA - if _SUPPORTS_TMA is None: - _SUPPORTS_TMA = torch.cuda.get_device_capability()[0] >= 9 return _SUPPORTS_TMA +# Helper to support allow_in_graph +try: + from torch.compiler import allow_in_graph +except ImportError: + from torch._dynamo import allow_in_graph + + +# Helper to detect if we're in tracing/compilation mode +def _is_tracing(*tensors): + """ + Check if tensors are fake tensors used during torch.compile tracing. + During tracing, tensors are FakeTensor/FunctionalTensor and we can't run Triton kernels. + During execution, tensors are real Tensors and we MUST run the kernels. + + NOTE: We do NOT use torch.compiler.is_compiling() because it returns True + during both tracing AND execution. We only want to skip kernels during tracing + when tensors are actually fake. + """ + for t in tensors: + name = type(t).__name__ + if name in ("FakeTensor", "FunctionalTensor", "FunctionalTensorWrapper"): + return True + return False + + _per_device_alloc_fns = {} @@ -83,6 +123,7 @@ def log_kernel_info( logger.debug(f"{kernel_name} autotuned best_config: {best_config}") +@allow_in_graph def grouped_gemm_forward( X: torch.Tensor, W: torch.Tensor, @@ -158,11 +199,21 @@ def grouped_gemm_forward( use_tma_store = False if use_tma or autotune: + # Respect global persistent allocator if set + if _HAS_SET_ALLOCATOR and not getattr(triton, "_unsloth_allocator_set", False): - def alloc_fn(size: int, alignment: int, stream: int): - return torch.empty(size, device = "cuda", dtype = torch.int8) + def alloc_fn(size: int, alignment: int, stream: int): + return torch.empty(size, device = "cuda", dtype = torch.int8) - triton.set_allocator(alloc_fn) + triton.set_allocator(alloc_fn) + + if W.ndim == 3: + num_experts = W.shape[0] + N = W.shape[1] + # K = W.shape[2] + else: + num_experts = m_sizes.shape[0] + N = W.shape[0] // num_experts X = X.view(-1, X.shape[-1]) W = W.view(-1, W.shape[-1]) @@ -188,9 +239,7 @@ def grouped_gemm_forward( total_tokens = X.shape[0] num_tokens = total_tokens // topk - num_experts = m_sizes.shape[0] _, K = X.shape - N = W.shape[0] // num_experts assert K == W.shape[1], f"K ({K}) must match W.shape[1] ({W.shape[1]})" if fuse_mul_post: @@ -212,8 +261,8 @@ def grouped_gemm_forward( ) y = torch.empty((total_tokens, N), device = X.device, dtype = X.dtype) - if total_tokens == 0 or N == 0: - return y + # if total_tokens == 0 or N == 0: + # return y NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count @@ -221,9 +270,9 @@ def grouped_gemm_forward( return (NUM_SMS,) if not autotune: - BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) - BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) - BLOCK_SIZE_M = min(total_tokens, BLOCK_SIZE_M) + # BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) + # BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) + pass if debug: print( @@ -276,16 +325,19 @@ def grouped_gemm_forward( if autotune else _grouped_gemm_forward_kernel ) - compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args) - if autotune: - log_kernel_info(compiled_kernel, kernel.best_config) - else: - log_kernel_info(compiled_kernel) + is_fake = _is_tracing(X, W) + if not is_fake: + compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args) + if autotune: + log_kernel_info(compiled_kernel, kernel.best_config) + else: + log_kernel_info(compiled_kernel) return y +@allow_in_graph def grouped_gemm_dX( dY: torch.Tensor, W: torch.Tensor, @@ -354,20 +406,28 @@ def grouped_gemm_dX( use_tma_store = False if use_tma or autotune: + # Respect global persistent allocator if set + if _HAS_SET_ALLOCATOR and not getattr(triton, "_unsloth_allocator_set", False): - def alloc_fn(size: int, alignment: int, stream: int): - # print(f"DEBUG::GROUPED_GEMM alloc_fn {size=} {alignment=} {stream=}") - return torch.empty(size, device = "cuda", dtype = torch.int8) + def alloc_fn(size: int, alignment: int, stream: int): + # print(f"DEBUG::GROUPED_GEMM alloc_fn {size=} {alignment=} {stream=}") + return torch.empty(size, device = "cuda", dtype = torch.int8) - triton.set_allocator(alloc_fn) + triton.set_allocator(alloc_fn) + + if W.ndim == 3: + num_experts = W.shape[0] + N = W.shape[1] + else: + num_experts = m_sizes.shape[0] + N = W.shape[0] // num_experts - num_experts = m_sizes.shape[0] dY = dY.view(-1, dY.shape[-1]) W = W.view(-1, W.shape[-1]) M_total, N_grad = dY.shape N_total, K = W.shape - N = N_total // num_experts + # N = N_total // num_experts assert N_grad == N, f"Grad_output N ({N_grad}) must match weight N ({N})" assert ( @@ -393,9 +453,9 @@ def grouped_gemm_dX( return (NUM_SMS,) if not autotune: - BLOCK_SIZE_M = min(M_total, BLOCK_SIZE_M) - BLOCK_SIZE_N = min(N_grad, BLOCK_SIZE_N) - BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) + # BLOCK_SIZE_N = min(N_grad, BLOCK_SIZE_N) + # BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) + pass if debug: print( @@ -437,15 +497,19 @@ def grouped_gemm_dX( } ) kernel = _autotuned_grouped_gemm_dX_kernel if autotune else _grouped_gemm_dX_kernel - compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args) - if autotune: - log_kernel_info(compiled_kernel, kernel.best_config) - else: - log_kernel_info(compiled_kernel) + is_fake = _is_tracing(dY, W) + if not is_fake: + compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args) + + if autotune: + log_kernel_info(compiled_kernel, kernel.best_config) + else: + log_kernel_info(compiled_kernel) return dX +@allow_in_graph def grouped_gemm_dW( X: torch.Tensor, dY: torch.Tensor, @@ -510,11 +574,13 @@ def grouped_gemm_dW( use_tma_store = False if use_tma or autotune: + # Respect global persistent allocator if set + if _HAS_SET_ALLOCATOR and not getattr(triton, "_unsloth_allocator_set", False): - def alloc_fn(size: int, alignment: int, stream: int): - return torch.empty(size, device = "cuda", dtype = torch.int8) + def alloc_fn(size: int, alignment: int, stream: int): + return torch.empty(size, device = "cuda", dtype = torch.int8) - triton.set_allocator(alloc_fn) + triton.set_allocator(alloc_fn) if permute_x or permute_y: assert gather_indices is not None @@ -541,9 +607,9 @@ def grouped_gemm_dW( dW = torch.zeros((num_experts, N, K), device = X.device, dtype = X.dtype) if not autotune: - BLOCK_SIZE_M = min(total_tokens, BLOCK_SIZE_M) - BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) - BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) + # BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) + # BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) + pass def grid(META): return (NUM_SMS,) @@ -607,12 +673,15 @@ def grouped_gemm_dW( ) kernel = _autotuned_grouped_gemm_dW_kernel if autotune else _grouped_gemm_dW_kernel - compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args) - if autotune: - log_kernel_info(compiled_kernel, kernel.best_config) - else: - log_kernel_info(compiled_kernel) + is_fake = _is_tracing(X, dY) + if not is_fake: + compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args) + + if autotune: + log_kernel_info(compiled_kernel, kernel.best_config) + else: + log_kernel_info(compiled_kernel) return dW @@ -680,6 +749,7 @@ class GroupedGemm(torch.autograd.Function): @staticmethod def backward(ctx, dY): + dY = dY.contiguous() X, W, m_sizes, gather_indices = ctx.saved_tensors topk = ctx.topk permute_x = ctx.permute_x diff --git a/unsloth/kernels/moe/grouped_gemm/kernels/autotuning.py b/unsloth/kernels/moe/grouped_gemm/kernels/autotuning.py index a185b5fd37..d25913975e 100644 --- a/unsloth/kernels/moe/grouped_gemm/kernels/autotuning.py +++ b/unsloth/kernels/moe/grouped_gemm/kernels/autotuning.py @@ -1,5 +1,18 @@ -# SPDX-License-Identifier: GNU Affero General Public License v3.0 -# Copyright 2023-present the Unsloth team. All rights reserved. +# Unsloth +# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as published +# by the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . """ Autotuning utils @@ -36,17 +49,39 @@ def convert_args_to_list(args): return [val_to_list(arg) for arg in args] +def _triton_supports_tma(): + """Check if current Triton version supports TMA API.""" + import triton.language as tl + + # Check for both old experimental and new stable API names + return hasattr(tl, "make_tensor_descriptor") or hasattr( + tl, "_experimental_make_tensor_descriptor" + ) + + +# Precompute at module import +# NOTE: TMA is disabled for now due to compatibility issues with permute_x/permute_y settings +# in the MoE grouped GEMM forward/backward passes. Re-enable once these are resolved. +_TRITON_HAS_TMA = False # _triton_supports_tma() + + def get_forward_configs( BLOCK_M = DEFAULT_M_BLOCK_SIZES, BLOCK_N = DEFAULT_N_BLOCK_SIZES, BLOCK_K = DEFAULT_K_BLOCK_SIZES, - TMA_LOAD_X = True, - TMA_LOAD_W = True, + TMA_LOAD_X = None, # Auto-detect if not specified + TMA_LOAD_W = None, # Auto-detect if not specified TMA_STORE = False, # NOTE: TMA_STORE is disabled for now num_warps = DEFAULT_NUM_WARPS, num_stages = DEFAULT_NUM_STAGES, num_ctas = DEFAULT_NUM_CTAS, ): + # Auto-detect TMA support + if TMA_LOAD_X is None: + TMA_LOAD_X = _TRITON_HAS_TMA + if TMA_LOAD_W is None: + TMA_LOAD_W = _TRITON_HAS_TMA + ( BLOCK_M, BLOCK_N, @@ -115,13 +150,18 @@ def get_dX_kernel_configs( BLOCK_M = DEFAULT_M_BLOCK_SIZES, BLOCK_N = DEFAULT_N_BLOCK_SIZES, BLOCK_K = DEFAULT_K_BLOCK_SIZES, - TMA_LOAD_dY = True, - TMA_LOAD_W = True, + TMA_LOAD_dY = None, # Auto-detect if not specified + TMA_LOAD_W = None, # Auto-detect if not specified TMA_STORE = False, # NOTE: TMA_STORE is disabled for now num_warps = DEFAULT_NUM_WARPS, num_stages = DEFAULT_NUM_STAGES, num_ctas = DEFAULT_NUM_CTAS, ): + # Auto-detect TMA support + if TMA_LOAD_dY is None: + TMA_LOAD_dY = _TRITON_HAS_TMA + if TMA_LOAD_W is None: + TMA_LOAD_W = _TRITON_HAS_TMA ( BLOCK_M, BLOCK_N, @@ -193,10 +233,15 @@ def get_dW_kernel_configs( num_warps = DEFAULT_NUM_WARPS, num_stages = DEFAULT_NUM_STAGES, num_ctas = DEFAULT_NUM_CTAS, - TMA_LOAD_dY = True, - TMA_LOAD_X = True, + TMA_LOAD_dY = None, # Auto-detect if not specified + TMA_LOAD_X = None, # Auto-detect if not specified TMA_STORE = False, ): + # Auto-detect TMA support + if TMA_LOAD_dY is None: + TMA_LOAD_dY = _TRITON_HAS_TMA + if TMA_LOAD_X is None: + TMA_LOAD_X = _TRITON_HAS_TMA ( BLOCK_M, BLOCK_N, @@ -291,8 +336,8 @@ def exceeds_smem_capacity( def common_prune_criteria(config: triton.Config, kwargs: dict, dtype): - from grouped_gemm.interface import supports_tma - from grouped_gemm.kernels.tuning import get_device_properties + from ..interface import supports_tma + from .tuning import get_device_properties smem_size = get_device_properties().SIZE_SMEM @@ -323,7 +368,7 @@ def common_prune_criteria(config: triton.Config, kwargs: dict, dtype): def maybe_disable_tma(config: triton.Config): - from grouped_gemm.interface import supports_tma + from ..interface import supports_tma tma_keys = [k for k in config.kwargs.keys() if k.startswith("USE_TMA_")] if not supports_tma(): diff --git a/unsloth/kernels/moe/grouped_gemm/kernels/backward.py b/unsloth/kernels/moe/grouped_gemm/kernels/backward.py index d8bdcb57e2..5e07056b52 100644 --- a/unsloth/kernels/moe/grouped_gemm/kernels/backward.py +++ b/unsloth/kernels/moe/grouped_gemm/kernels/backward.py @@ -5,7 +5,7 @@ import torch import triton import triton.language as tl -from grouped_gemm.kernels.autotuning import ( +from .autotuning import ( get_dW_kernel_configs, get_dX_kernel_configs, prune_dX_configs, @@ -53,11 +53,11 @@ def _grouped_gemm_dX_kernel( m_sizes_ptr, # problem sizes NUM_EXPERTS: tl.constexpr, - NUM_TOKENS: tl.constexpr, + NUM_TOKENS, TOPK: tl.constexpr, N: tl.constexpr, K: tl.constexpr, - NUM_SMS: tl.constexpr, + NUM_SMS, # Tuning parameters BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, @@ -69,7 +69,7 @@ def _grouped_gemm_dX_kernel( USE_TMA_STORE: tl.constexpr = False, FLATTEN: tl.constexpr = True, ) -> None: - TOTAL_TOKENS: tl.constexpr = NUM_TOKENS * TOPK + TOTAL_TOKENS = NUM_TOKENS * TOPK output_dtype = dX_ptr.dtype.element_ty tidx = tl.program_id(0) @@ -82,7 +82,7 @@ def _grouped_gemm_dX_kernel( # Also, we are defining a single global descriptor with single block shape # Need to check that this does not result in errors when crossing expert boundaries if USE_TMA_LOAD_dY: - dY_desc = tl._experimental_make_tensor_descriptor( + dY_desc = tl.make_tensor_descriptor( dY_ptr, shape = [TOTAL_TOKENS, N], strides = [N, 1], @@ -91,7 +91,7 @@ def _grouped_gemm_dX_kernel( if USE_TMA_LOAD_W: expert_stride = N * K - w_desc = tl._experimental_make_tensor_descriptor( + w_desc = tl.make_tensor_descriptor( w_ptr, shape = [NUM_EXPERTS, N, K], strides = [expert_stride, K, 1], @@ -123,7 +123,7 @@ def _grouped_gemm_dX_kernel( tl.static_assert( K % BLOCK_SIZE_K == 0, "K must be divisible by BLOCK_SIZE_K" ) - dX_desc = tl._experimental_make_tensor_descriptor( + dX_desc = tl.make_tensor_descriptor( dX_ptr, shape = [m_end, K], strides = [K, 1], @@ -232,6 +232,7 @@ def _grouped_gemm_dX_kernel( # TODO: check if predication along K is needed since we checked that K is divisible by BLOCK_SIZE_K in the forward kernel # [M, N] @ [N, K] -> [M, K] + dY = dY.to(w.dtype) accumulator += tl.dot(dY, w) # NOTE: no transpose of b # Advance A along contiguous dimension @@ -266,7 +267,8 @@ def _grouped_gemm_dX_kernel( _autotuned_grouped_gemm_dX_kernel = triton.autotune( configs = get_dX_kernel_configs(), prune_configs_by = {"early_config_prune": prune_dX_configs}, - key = ["NUM_EXPERTS", "NUM_TOKENS", "N", "K", "PERMUTE_X", "PERMUTE_Y"], + # NOTE: NUM_TOKENS removed from key to avoid recompilation for every sequence length + key = ["NUM_EXPERTS", "N", "K", "PERMUTE_X", "PERMUTE_Y"], )(_grouped_gemm_dX_kernel) """ @@ -298,12 +300,12 @@ def _grouped_gemm_dW_kernel( m_sizes_ptr, gather_indices_ptr, # problem sizes - NUM_TOKENS: tl.constexpr, + NUM_TOKENS, TOPK: tl.constexpr, NUM_EXPERTS: tl.constexpr, N: tl.constexpr, K: tl.constexpr, - NUM_SMS: tl.constexpr, + NUM_SMS, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, BLOCK_SIZE_M: tl.constexpr, @@ -315,14 +317,14 @@ def _grouped_gemm_dW_kernel( FLATTEN: tl.constexpr = True, acc_dtype: tl.constexpr = tl.float32, ) -> None: - TOTAL_TOKENS: tl.constexpr = NUM_TOKENS * TOPK + TOTAL_TOKENS = NUM_TOKENS * TOPK TMA_LOAD_BOTH: tl.constexpr = USE_TMA_LOAD_X and USE_TMA_LOAD_dY tidx = tl.program_id(0) output_dtype = dW_ptr.dtype.element_ty if USE_TMA_LOAD_dY and not TMA_LOAD_BOTH: - dY_desc = tl._experimental_make_tensor_descriptor( + dY_desc = tl.make_tensor_descriptor( dY_ptr, shape = [TOTAL_TOKENS, N], strides = [N, 1], @@ -330,7 +332,7 @@ def _grouped_gemm_dW_kernel( ) if USE_TMA_LOAD_X and not TMA_LOAD_BOTH: - x_desc = tl._experimental_make_tensor_descriptor( + x_desc = tl.make_tensor_descriptor( x_ptr, shape = [TOTAL_TOKENS, K], strides = [K, 1], @@ -349,7 +351,7 @@ def _grouped_gemm_dW_kernel( if USE_TMA_STORE: tl.static_assert(N % BLOCK_SIZE_N == 0, "N must be divisible by BLOCK_SIZE_N") tl.static_assert(K % BLOCK_SIZE_K == 0, "K must be divisible by BLOCK_SIZE_K") - dW_desc = tl._experimental_make_tensor_descriptor( + dW_desc = tl.make_tensor_descriptor( dW_ptr, shape = [NUM_EXPERTS, N, K], strides = [N * K, K, 1], @@ -390,14 +392,14 @@ def _grouped_gemm_dW_kernel( if m_size > 0: if TMA_LOAD_BOTH: - dY_desc = tl._experimental_make_tensor_descriptor( + dY_desc = tl.make_tensor_descriptor( dY_ptr, shape = [m_end, N], strides = [N, 1], block_shape = [BLOCK_SIZE_M, BLOCK_SIZE_N], ) - x_desc = tl._experimental_make_tensor_descriptor( + x_desc = tl.make_tensor_descriptor( x_ptr, shape = [m_end, K], strides = [K, 1], @@ -475,7 +477,7 @@ def _grouped_gemm_dW_kernel( ) accumulator += tl.dot( - dY.T, # [BLOCK_N, BLOCK_M] + dY.T.to(x.dtype), # [BLOCK_N, BLOCK_M] x, # [BLOCK_M, BLOCK_K] ) @@ -498,5 +500,6 @@ def _grouped_gemm_dW_kernel( _autotuned_grouped_gemm_dW_kernel = triton.autotune( configs = get_dW_kernel_configs(), prune_configs_by = {"early_config_prune": prune_kernel_configs_backward_dW}, - key = ["NUM_EXPERTS", "NUM_TOKENS", "N", "K", "PERMUTE_X", "PERMUTE_Y"], + # NOTE: NUM_TOKENS removed from key to avoid recompilation for every sequence length + key = ["NUM_EXPERTS", "N", "K", "PERMUTE_X", "PERMUTE_Y"], )(_grouped_gemm_dW_kernel) diff --git a/unsloth/kernels/moe/grouped_gemm/kernels/forward.py b/unsloth/kernels/moe/grouped_gemm/kernels/forward.py index f84694e91e..a42ec5ffe9 100644 --- a/unsloth/kernels/moe/grouped_gemm/kernels/forward.py +++ b/unsloth/kernels/moe/grouped_gemm/kernels/forward.py @@ -5,7 +5,7 @@ import torch import triton import triton.language as tl -from grouped_gemm.kernels.autotuning import ( +from .autotuning import ( get_forward_configs, prune_kernel_configs_fwd, ) @@ -31,11 +31,11 @@ def _grouped_gemm_forward_kernel( topk_weights_ptr, # Constant problem shapes NUM_EXPERTS: tl.constexpr, - NUM_TOKENS: tl.constexpr, + NUM_TOKENS, TOPK: tl.constexpr, N: tl.constexpr, K: tl.constexpr, - NUM_SMS: tl.constexpr, + NUM_SMS, # Tuning params BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, @@ -53,7 +53,7 @@ def _grouped_gemm_forward_kernel( ) -> None: tl.static_assert(K % BLOCK_SIZE_K == 0) - TOTAL_TOKENS: tl.constexpr = NUM_TOKENS * TOPK + TOTAL_TOKENS = NUM_TOKENS * TOPK SHOULD_PERMUTE: tl.constexpr = PERMUTE_X or PERMUTE_Y SHOULD_FUSE_MUL: tl.constexpr = FUSE_MUL_PRE or FUSE_MUL_POST SHOULD_PERMUTE_OR_FUSE: tl.constexpr = SHOULD_PERMUTE or SHOULD_FUSE_MUL @@ -66,7 +66,7 @@ def _grouped_gemm_forward_kernel( # Also, we are defining a single global descriptor with single block shape # Need to check that this does not result in errors when crossing expert boundaries if USE_TMA_LOAD_X: - x_desc = tl._experimental_make_tensor_descriptor( + x_desc = tl.make_tensor_descriptor( x_ptr, shape = [TOTAL_TOKENS, K], strides = [K, 1], @@ -75,7 +75,7 @@ def _grouped_gemm_forward_kernel( if USE_TMA_LOAD_W: expert_stride = N * K - w_desc = tl._experimental_make_tensor_descriptor( + w_desc = tl.make_tensor_descriptor( w_ptr, shape = [NUM_EXPERTS, N, K], strides = [expert_stride, K, 1], @@ -100,7 +100,7 @@ def _grouped_gemm_forward_kernel( # Need to create tma_store within loop since we need to predicate stores based on m_size if USE_TMA_STORE: - y_desc = tl._experimental_make_tensor_descriptor( + y_desc = tl.make_tensor_descriptor( y_ptr, # + m_start * N, shape = [m_end, N], strides = [N, 1], @@ -213,6 +213,7 @@ def _grouped_gemm_forward_kernel( ) w = tl.reshape(w, (BLOCK_SIZE_N, BLOCK_SIZE_K)) + x = x.to(w.dtype) accumulator += tl.dot(x, w.T) if not USE_TMA_LOAD_X: @@ -253,9 +254,10 @@ def _grouped_gemm_forward_kernel( _autotuned_grouped_gemm_forward_kernel = triton.autotune( configs = get_forward_configs(), prune_configs_by = {"early_config_prune": prune_kernel_configs_fwd}, + # NOTE: NUM_TOKENS removed from key to avoid recompilation for every sequence length + # The kernel handles variable token counts via m_sizes and tile-based processing key = [ "NUM_EXPERTS", - "NUM_TOKENS", "N", "K", "PERMUTE_X", diff --git a/unsloth/kernels/moe/grouped_gemm/kernels/tuning.py b/unsloth/kernels/moe/grouped_gemm/kernels/tuning.py index 1f641478b4..00d4824703 100644 --- a/unsloth/kernels/moe/grouped_gemm/kernels/tuning.py +++ b/unsloth/kernels/moe/grouped_gemm/kernels/tuning.py @@ -15,7 +15,7 @@ import torch import triton from triton.runtime.errors import OutOfResources -from grouped_gemm.kernels.autotuning import ( +from .autotuning import ( BOOLS, DEFAULT_K_BLOCK_SIZES, DEFAULT_M_BLOCK_SIZES, diff --git a/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py b/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py index 4010c77ce4..6bb0bfb0c3 100644 --- a/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py +++ b/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py @@ -9,13 +9,13 @@ import torch.nn.functional as F from transformers.models.llama4 import Llama4TextConfig from transformers.models.llama4.modeling_llama4 import Llama4TextMoe -from grouped_gemm.interface import grouped_gemm -from grouped_gemm.kernels.tuning import ( +from ...interface import grouped_gemm +from ...kernels.tuning import ( KernelConfigBackward_dW, KernelConfigBackward_dX, KernelConfigForward, ) -from grouped_gemm.reference.moe_ops import ( +from ..moe_ops import ( get_routing_indices, permute, torch_grouped_gemm, diff --git a/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py index ace6a7714a..31c635ba37 100644 --- a/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py +++ b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py @@ -12,13 +12,13 @@ from transformers.models.qwen3_moe.modeling_qwen3_moe import ( Qwen3MoeSparseMoeBlock, ) -from grouped_gemm.interface import grouped_gemm -from grouped_gemm.kernels.tuning import ( +from ...interface import grouped_gemm +from ...kernels.tuning import ( KernelConfigBackward_dW, KernelConfigBackward_dX, KernelConfigForward, ) -from grouped_gemm.reference.moe_ops import ( +from ..moe_ops import ( get_routing_indices, permute, torch_grouped_gemm, diff --git a/unsloth/kernels/moe/grouped_gemm/reference/moe_block.py b/unsloth/kernels/moe/grouped_gemm/reference/moe_block.py index 0d497f380c..2a015252c6 100644 --- a/unsloth/kernels/moe/grouped_gemm/reference/moe_block.py +++ b/unsloth/kernels/moe/grouped_gemm/reference/moe_block.py @@ -5,13 +5,13 @@ import torch from transformers.models.qwen3_moe.configuration_qwen3_moe import Qwen3MoeConfig from transformers.models.qwen3_moe.modeling_qwen3_moe import Qwen3MoeSparseMoeBlock -from grouped_gemm.interface import grouped_gemm -from grouped_gemm.kernels.tuning import ( +from ..interface import grouped_gemm +from ..kernels.tuning import ( KernelConfigBackward_dW, KernelConfigBackward_dX, KernelConfigForward, ) -from grouped_gemm.reference.moe_ops import ( +from .moe_ops import ( Qwen3MoeGroupedGEMMBlock, permute, unpermute, diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 76952b00a5..48e5683076 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -__version__ = "2026.1.4" +__version__ = "2026.2.1" __all__ = [ "SUPPORTS_BFLOAT16", @@ -59,9 +59,11 @@ __all__ = [ "unsloth_fused_ce_loss", "patch_unsloth_smart_gradient_checkpointing", "unpatch_unsloth_smart_gradient_checkpointing", + "apply_unsloth_gradient_checkpointing", "patch_compiled_autograd", "process_vision_info", "unsloth_compile_transformers", + "prefer_flex_attn_if_supported", "patch_fast_lora", "validate_loftq_config", "RaiseUninitialized", @@ -73,6 +75,8 @@ __all__ = [ "verify_fp8_support_if_applicable", "_get_inference_mode_context_manager", "hf_login", + "is_moe_model", + "get_moe_target_parameters", "make_fast_generate_wrapper", ] @@ -148,6 +152,68 @@ from unsloth_zoo.temporary_patches import ( TEMPORARY_PATCHES, ) + +def apply_unsloth_gradient_checkpointing( + use_gradient_checkpointing, max_seq_length, dtype +): + """ + Apply gradient checkpointing with smart heuristics. + + For seq < 512, the overhead of gradient offloading in gc="unsloth" mode + is not worth it. Benchmarks show standard gc is faster for small sequences. + + Args: + use_gradient_checkpointing: "unsloth", True, False, or None + max_seq_length: The maximum sequence length + dtype: The model dtype for patching + + Returns: + The effective use_gradient_checkpointing value (may change from "unsloth" to True) + """ + if use_gradient_checkpointing == "unsloth": + # Gradient offloading overhead is not worth it for small sequences. + # Benchmarks show crossover point is around seq_len 384-512. + # For seq < 512, standard gradient checkpointing is faster. + if max_seq_length < 512: + unpatch_unsloth_smart_gradient_checkpointing() + return True + else: + patch_unsloth_smart_gradient_checkpointing(dtype = dtype) + return "unsloth" + elif use_gradient_checkpointing in (True, False): + # User explicitly set True or False - unpatch any previous "unsloth" patching + unpatch_unsloth_smart_gradient_checkpointing() + return use_gradient_checkpointing + return use_gradient_checkpointing + + +def prefer_flex_attn_if_supported(model_class, config): + if os.environ.get("UNSLOTH_ENABLE_FLEX_ATTENTION", "1") == "0": + return None + try: + from transformers.utils.import_utils import is_torch_flex_attn_available + + if not is_torch_flex_attn_available(): + return None + if model_class is None or not getattr( + model_class, "_supports_flex_attn", False + ): + return None + # GPT-OSS uses eager attention during inference since flex attention + # returns incorrect results (likely due to left padding issues). + # Skip setting flex_attention to avoid BlockMask type errors. + model_type = getattr(config, "model_type", "") if config else "" + if model_type == "gpt_oss": + return None + if config is not None: + setattr(config, "_attn_implementation", "flex_attention") + if hasattr(config, "attn_implementation"): + setattr(config, "attn_implementation", "flex_attention") + return "flex_attention" + except Exception: + return None + + for temporary_patch in TEMPORARY_PATCHES: temporary_patch() @@ -535,6 +601,12 @@ try: from transformers.configuration_utils import layer_type_validation except: pass + +try: + # Transformers 5.0+ uses RotaryEmbeddingConfigMixin as a base class for configs + from transformers.modeling_rope_utils import RotaryEmbeddingConfigMixin +except: + pass from transformers import __version__ as transformers_version try: @@ -1088,8 +1160,12 @@ def has_internet(host = "8.8.8.8", port = 53, timeout = 3): return False try: socket.setdefaulttimeout(timeout) - socket.socket(socket.AF_INET, socket.SOCK_STREAM).connect((host, port)) - return True + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + try: + sock.connect((host, port)) + return True + finally: + sock.close() except socket.error as ex: return False @@ -2428,6 +2504,117 @@ def hf_login(token: Optional[str] = None) -> Optional[str]: return token +# ============================================= +# MoE (Mixture of Experts) Detection and LoRA Utilities + + +def is_moe_model(model) -> bool: + """ + Detect if a model is a Mixture of Experts (MoE) model. + + Args: + model: The model to check (can be HF model or config) + + Returns: + True if the model is an MoE model, False otherwise + """ + config = getattr(model, "config", model) + + # Different MoE models use different config attribute names: + # - Qwen3-MoE: num_experts + # - GLM4-MoE: n_routed_experts, num_local_experts + # - Mixtral: num_local_experts + num_experts = None + for attr in ("num_experts", "n_routed_experts", "num_local_experts"): + num_experts = getattr(config, attr, None) + if num_experts is not None: + break + + # Check text_config for VL models + if num_experts is None and hasattr(config, "text_config"): + for attr in ("num_experts", "n_routed_experts", "num_local_experts"): + num_experts = getattr(config.text_config, attr, None) + if num_experts is not None: + break + + return num_experts is not None and num_experts > 0 + + +def get_moe_target_parameters(model, target_modules = None) -> Optional[List[str]]: + """ + Get the target_parameters for MoE expert layers if applicable. + + For MoE models, returns the parameter paths for expert weights + (gate_up_proj, down_proj) that should be targeted by PEFT's + target_parameters for LoRA on nn.Parameter. + + Only includes MoE parameters that match what's in target_modules: + - If "down_proj" is in target_modules -> includes "mlp.experts.down_proj" + - If "gate_proj" or "up_proj" is in target_modules -> includes "mlp.experts.gate_up_proj" + + Args: + model: The model to get target parameters for + target_modules: List/tuple of target module names to match against + + Returns: + List of parameter paths for MoE experts, or None if not an MoE model + """ + if not is_moe_model(model): + return None + + config = getattr(model, "config", model) + # Get num_experts from various possible config attributes + num_experts = None + for attr in ("num_experts", "n_routed_experts", "num_local_experts"): + num_experts = getattr(config, attr, None) + if num_experts is not None: + break + if num_experts is None and hasattr(config, "text_config"): + for attr in ("num_experts", "n_routed_experts", "num_local_experts"): + num_experts = getattr(config.text_config, attr, None) + if num_experts is not None: + break + if num_experts is None: + num_experts = 0 + + # Determine which MoE parameters to include based on target_modules + moe_params = [] + + # Normalize target_modules to a set for efficient lookup + if target_modules is None: + # If no target_modules specified, include all MoE params + target_set = {"gate_proj", "up_proj", "down_proj", "gate_up_proj"} + elif isinstance(target_modules, str): + target_set = {target_modules} + # Heuristic for regex matching MLPs + if "proj" in target_modules and ( + "mlp" in target_modules or "ffn" in target_modules + ): + target_set.update({"gate_proj", "up_proj", "down_proj", "gate_up_proj"}) + else: + target_set = set(target_modules) if target_modules else set() + + # gate_up_proj combines both gate_proj and up_proj in MoE + # Also match "gate_up_proj" directly since users may specify the fused name + if ( + "gate_proj" in target_set + or "up_proj" in target_set + or "gate_up_proj" in target_set + ): + moe_params.append("mlp.experts.gate_up_proj") + + if "down_proj" in target_set: + moe_params.append("mlp.experts.down_proj") + + if moe_params: + print( + f"Unsloth: Detected MoE model with {num_experts} experts - enabling LoRA on: {moe_params}" + ) + return moe_params + + return None + + def make_fast_generate_wrapper(original_generate): """ Creates a wrapper around model.generate that checks for incorrect diff --git a/unsloth/models/glm4_moe.py b/unsloth/models/glm4_moe.py new file mode 100644 index 0000000000..5d04b2f1d0 --- /dev/null +++ b/unsloth/models/glm4_moe.py @@ -0,0 +1,450 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +GLM-4.7 Flash (GLM4 MoE Lite) optimized implementation using grouped GEMM. + +Key architecture differences from Qwen3 MoE: +- Router uses sigmoid activation (not softmax) +- Has routed_scaling_factor of 1.8 +- Has 1 shared expert that processes all tokens +- Uses group-based selection before topk +- Uses MLA (Multi-head Latent Attention) +""" + +from .llama import * +import os +from ._utils import __version__ +from .llama import ( + LlamaRotaryEmbedding, + LlamaLinearScalingRotaryEmbedding, + fix_prepare_inputs_for_generation, + fast_rms_layernorm_inference, + fast_swiglu_inference, + LlamaModel_fast_forward, + LlamaModel_fast_forward_inference, + CausalLM_fast_forward, + PeftModel_fast_forward, +) +import torch +import torch.nn.functional as F +from typing import Optional, Tuple +from ..kernels import fast_rms_layernorm + +# Import the grouped gemm utilities from unsloth kernels +# The grouped_gemm module expects its parent directory to be in sys.path +HAS_GROUPED_GEMM = False +try: + import sys + import os + + # Add the moe directory (parent of grouped_gemm) to sys.path + _moe_path = os.path.join( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "kernels", "moe" + ) + if _moe_path not in sys.path: + sys.path.insert(0, _moe_path) + + # Import grouped_gemm package first to apply TMA compatibility shim + # This patches triton.language to support both old and new TMA API names + import grouped_gemm # noqa: F401 - triggers TMA compatibility shim + + from grouped_gemm.interface import grouped_gemm + from grouped_gemm.reference.moe_ops import ( + get_routing_indices, + permute, + unpermute, + ) + + HAS_GROUPED_GEMM = True +except ImportError as e: + import warnings + + warnings.warn( + f"Grouped GEMM not available: {e}. MoE will use fallback implementation." + ) + + +# Import transformers GLM4 MoE Lite classes +try: + from transformers.models.glm4_moe_lite.modeling_glm4_moe_lite import ( + Glm4MoeLiteAttention, + Glm4MoeLiteMoE, + Glm4MoeLiteMLP, + Glm4MoeLiteNaiveMoe, + Glm4MoeLiteTopkRouter, + Glm4MoeLiteDecoderLayer, + Glm4MoeLiteModel, + Glm4MoeLiteForCausalLM, + Glm4MoeLiteRMSNorm, + ) + + HAS_GLM4_MOE = True +except ImportError: + HAS_GLM4_MOE = False + + # Create dummy classes for type checking + class Glm4MoeLiteAttention: + pass + + class Glm4MoeLiteMoE: + pass + + class Glm4MoeLiteMLP: + pass + + class Glm4MoeLiteNaiveMoe: + pass + + class Glm4MoeLiteTopkRouter: + pass + + class Glm4MoeLiteDecoderLayer: + pass + + class Glm4MoeLiteModel: + pass + + class Glm4MoeLiteForCausalLM: + pass + + +torch_nn_functional_silu = torch.nn.functional.silu + + +def Glm4MoeLiteMoE_fast_forward(self, hidden_states): + """ + Optimized MoE forward pass using grouped GEMM. + + GLM4 MoE specifics: + - Uses sigmoid router activation (not softmax) + - Has routed_scaling_factor of 1.8 + - Has 1 shared expert that always processes all tokens + - Uses group-based selection with topk_group + """ + residuals = hidden_states + orig_shape = hidden_states.shape + batch_size, seq_len, hidden_dim = orig_shape + num_tokens = batch_size * seq_len + + # Flatten hidden states for routing + hidden_states = hidden_states.view(-1, hidden_dim) + + # Router computation + router_logits = self.gate(hidden_states) # [num_tokens, n_routed_experts] + topk_indices, topk_weights = self.route_tokens_to_experts(router_logits) + # Cast routing weights to match hidden_states dtype (Qwen3 pattern) + # Sigmoid router returns fp32, but hidden_states may be bf16 + topk_weights = topk_weights.to(hidden_states.dtype) + + # Get routing indices for grouped GEMM + with torch.no_grad(): + token_counts_by_expert, gather_indices = get_routing_indices( + topk_indices, self.n_routed_experts + ) + + # Use grouped GEMM for expert computation + if HAS_GROUPED_GEMM: + # Cast hidden_states to match expert weights dtype + # Under autocast, hidden_states may be fp32 while weights are bf16 + hidden_states = hidden_states.to(self.experts.gate_up_proj.dtype) + + # First grouped GEMM: gate_up_proj with permute_x + # Input: [num_tokens, hidden_dim] -> Output: [total_tokens, 2*intermediate_dim] + intermediate = grouped_gemm( + X = hidden_states, + W = self.experts.gate_up_proj, + m_sizes = token_counts_by_expert.int(), + topk = self.top_k, + gather_indices = gather_indices, + permute_x = True, + permute_y = False, + autotune = True, + is_first_gemm = True, + ) + + # Activation: SiLU(gate) * up + gate, up = intermediate.chunk(2, dim = -1) + intermediate = torch_nn_functional_silu(gate) * up + + # Second grouped GEMM: down_proj with permute_y + # Input: [total_tokens, intermediate_dim] -> Output: [total_tokens, hidden_dim] + expert_output = grouped_gemm( + X = intermediate, + W = self.experts.down_proj, + m_sizes = token_counts_by_expert.int(), + topk = self.top_k, + gather_indices = gather_indices, + permute_x = False, + permute_y = True, + autotune = True, + is_first_gemm = False, + ) + + # Merge topk weights: [num_tokens, top_k, hidden_dim] -> [num_tokens, hidden_dim] + hidden_states = ( + expert_output.view(num_tokens, self.top_k, hidden_dim) + * topk_weights.unsqueeze(-1) + ).sum(dim = 1) + else: + # Fallback to naive implementation + hidden_states = self.experts(hidden_states, topk_indices, topk_weights) + + # Add shared expert output + hidden_states = hidden_states + self.shared_experts(residuals.view(-1, hidden_dim)) + + return hidden_states.view(*orig_shape) + + +def Glm4MoeLiteNaiveMoe_fast_forward( + self, + hidden_states: torch.Tensor, + top_k_index: torch.Tensor, + top_k_weights: torch.Tensor, +) -> torch.Tensor: + """ + Optimized expert forward using grouped GEMM. + + Args: + hidden_states: [num_tokens, hidden_dim] + top_k_index: [num_tokens, top_k] indices of selected experts + top_k_weights: [num_tokens, top_k] weights for selected experts + + Returns: + [num_tokens, hidden_dim] output after weighted sum of expert outputs + """ + num_tokens, hidden_dim = hidden_states.shape + top_k = top_k_index.shape[1] + # Cast routing weights to match hidden_states dtype (Qwen3 pattern) + top_k_weights = top_k_weights.to(hidden_states.dtype) + + if not HAS_GROUPED_GEMM: + # Fallback to original naive implementation + final_hidden_states = torch.zeros_like(hidden_states) + with torch.no_grad(): + expert_mask = torch.nn.functional.one_hot( + top_k_index, num_classes = self.num_experts + ) + expert_mask = expert_mask.permute(2, 1, 0) + expert_hit = torch.greater(expert_mask.sum(dim = (-1, -2)), 0).nonzero() + + for expert_idx in expert_hit: + expert_idx = expert_idx[0] + if expert_idx == self.num_experts: + continue + top_k_pos, token_idx = torch.where(expert_mask[expert_idx]) + current_state = hidden_states[token_idx] + gate, up = torch.nn.functional.linear( + current_state, self.gate_up_proj[expert_idx] + ).chunk(2, dim = -1) + current_hidden_states = self.act_fn(gate) * up + current_hidden_states = torch.nn.functional.linear( + current_hidden_states, self.down_proj[expert_idx] + ) + current_hidden_states = ( + current_hidden_states * top_k_weights[token_idx, top_k_pos, None] + ) + final_hidden_states.index_add_( + 0, token_idx, current_hidden_states.to(final_hidden_states.dtype) + ) + + return final_hidden_states + + # Get routing indices for grouped GEMM + with torch.no_grad(): + token_counts_by_expert, gather_indices = get_routing_indices( + top_k_index, self.num_experts + ) + + # Cast hidden_states to match expert weights dtype + # Under autocast, hidden_states may be fp32 while weights are bf16 + hidden_states = hidden_states.to(self.gate_up_proj.dtype) + + # First grouped GEMM: gate_up_proj + intermediate = grouped_gemm( + X = hidden_states, + W = self.gate_up_proj, + m_sizes = token_counts_by_expert.int(), + topk = top_k, + gather_indices = gather_indices, + permute_x = True, + permute_y = False, + autotune = True, + is_first_gemm = True, + ) + + # Activation: SiLU(gate) * up + gate, up = intermediate.chunk(2, dim = -1) + intermediate = self.act_fn(gate) * up + + # Second grouped GEMM: down_proj + expert_output = grouped_gemm( + X = intermediate, + W = self.down_proj, + m_sizes = token_counts_by_expert.int(), + topk = top_k, + gather_indices = gather_indices, + permute_x = False, + permute_y = True, + autotune = True, + is_first_gemm = False, + ) + + # Merge topk weights + final_hidden_states = ( + expert_output.view(num_tokens, top_k, hidden_dim) * top_k_weights.unsqueeze(-1) + ).sum(dim = 1) + + return final_hidden_states + + +def Glm4MoeLiteDecoderLayer_fast_forward( + self, + hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.LongTensor] = None, + past_key_values = None, + use_cache: bool = False, + cache_position: Optional[torch.LongTensor] = None, + position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + **kwargs, +) -> torch.Tensor: + """ + Optimized decoder layer forward with fast RMS layernorm. + """ + # Check if we're in inference mode + is_inference = use_cache and hasattr(self, "_flag_for_generation") + + if is_inference: + # Self-attention with fast inference path + residual = hidden_states + hidden_states = fast_rms_layernorm_inference( + self.input_layernorm, hidden_states + ) + hidden_states, _ = self.self_attn( + hidden_states = hidden_states, + attention_mask = attention_mask, + position_ids = position_ids, + past_key_values = past_key_values, + use_cache = use_cache, + cache_position = cache_position, + position_embeddings = position_embeddings, + **kwargs, + ) + hidden_states = residual + hidden_states + + # MLP/MoE + residual = hidden_states + hidden_states = fast_rms_layernorm_inference( + self.post_attention_layernorm, hidden_states + ) + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + else: + # Training path + residual = hidden_states + hidden_states = fast_rms_layernorm(self.input_layernorm, hidden_states) + hidden_states, _ = self.self_attn( + hidden_states = hidden_states, + attention_mask = attention_mask, + position_ids = position_ids, + past_key_values = past_key_values, + use_cache = use_cache, + cache_position = cache_position, + position_embeddings = position_embeddings, + **kwargs, + ) + hidden_states = residual + hidden_states + + # MLP/MoE + residual = hidden_states + hidden_states = fast_rms_layernorm(self.post_attention_layernorm, hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + + return hidden_states + + +def Glm4MoeLiteMLP_fast_forward(self, x): + """ + Optimized MLP forward using fused SwiGLU. + """ + return fast_swiglu_inference(self, x) + + +class FastGLM47Model(FastLlamaModel): + """ + Fast GLM-4.7 Flash (GLM4 MoE Lite) model with grouped GEMM optimization. + + This provides 2-3x throughput improvement for MoE layers by: + - Replacing sequential expert loops with grouped GEMM operations + - Fusing permutation operations into the GEMM kernels + - Using optimized RMS LayerNorm and SwiGLU implementations + """ + + @staticmethod + def pre_patch(): + if not HAS_GLM4_MOE: + raise ImportError( + "Unsloth: GLM4 MoE Lite support requires transformers >= 5.0.0. " + "Please upgrade with: pip install --upgrade transformers" + ) + + # Patch MoE forward with grouped GEMM optimization + # TMA compatibility is handled by grouped_gemm/__init__.py which patches + # triton.language to support both old (_experimental_make_tensor_descriptor) + # and new (make_tensor_descriptor) API names + if HAS_GROUPED_GEMM: + Glm4MoeLiteNaiveMoe.forward = Glm4MoeLiteNaiveMoe_fast_forward + Glm4MoeLiteMoE.forward = Glm4MoeLiteMoE_fast_forward + + # Note: We don't patch the following for GLM4 MoE because: + # - GLM4 uses MLA (Multi-head Latent Attention) which has different projection names + # - Glm4MoeLiteRotaryEmbedding doesn't have extend_rope_embedding method + # - The decoder layer and model forward functions assume Llama-compatible infrastructure + + return + + @staticmethod + def from_pretrained( + model_name = "unsloth/GLM-4.7-Flash", + 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, + ): + # Pop kwargs that are used by loader but not passed to model + kwargs.pop("unsloth_force_compile", None) + + return FastLlamaModel.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 = FastGLM47Model, + tokenizer_name = tokenizer_name, + trust_remote_code = trust_remote_code, + **kwargs, + ) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 39f2ba1460..f18a07ac3c 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -19,7 +19,7 @@ import functools from typing import Optional, Tuple, List, Union from ._utils import * -from ._utils import patch_unsloth_smart_gradient_checkpointing +from ._utils import apply_unsloth_gradient_checkpointing from ._utils import __version__, importlib_version from ._utils import move_to_device from ._utils import ( @@ -152,21 +152,22 @@ from peft.utils.other import ModulesToSaveWrapper def _offload_frozen_module_for_training( module: ModulesToSaveWrapper, device_type: str, - offload_device: str = "cpu", + offload_device: Optional[str] = "cpu", ) -> None: """ Offload frozen module to CPU and configure trainable copy for mixed precision training. This function optimizes memory usage by: 1. Moving the trainable copy to the target device with appropriate precision - 2. Offloading the original frozen module to CPU/disk to free VRAM + 2. Optionally offloading the original frozen module to CPU/disk to free VRAM 3. Converting float16 to float32 for compatibility with certain GPUs (e.g., Tesla T4) Args: module: The module to configure. Must be a ModulesToSaveWrapper with a `modules_to_save` attribute containing trainable and original modules. device_type: Target device string for training (e.g., "cuda:0", "xpu:0") - offload_device: Device to offload frozen parameters (default: "cpu") + offload_device: Device to offload frozen parameters (default: "cpu"). + If None, the original frozen module remains on its current device. Note: Currently only "cpu" is supported; disk offloading is planned. Returns: @@ -174,7 +175,7 @@ def _offload_frozen_module_for_training( Note: - Float16 weights are automatically promoted to float32 for GPU compatibility - - Original frozen parameters are moved to CPU to reduce active VRAM usage + - When offload_device is specified, frozen parameters are moved to free VRAM - Future versions will support disk-based offloading for even larger models See Also: @@ -196,7 +197,8 @@ def _offload_frozen_module_for_training( module.modules_to_save.default.requires_grad_(True) # [TODO] Move old module to CPU - should be disk! - module.original_module.to(device = offload_device, non_blocking = True) + if offload_device is not None: + module.original_module.to(device = offload_device, non_blocking = True) module.original_module.requires_grad_(False) @@ -2286,6 +2288,10 @@ class FastLlamaModel: model_function = MODEL_FOR_CAUSAL_LM_MAPPING[model_config.__class__] IS_FALCON_H1 = model_config.model_type.startswith("falcon_h1") + preferred_attn_impl = ( + prefer_flex_attn_if_supported(model_function, model_config) or "eager" + ) + has_rope_scaling = False try: with open(inspect.getfile(model_function), "r", encoding = "utf-8") as file: @@ -2364,7 +2370,7 @@ class FastLlamaModel: token = token, max_position_embeddings = max_position_embeddings, trust_remote_code = trust_remote_code, - attn_implementation = "eager", + attn_implementation = preferred_attn_impl, **kwargs, ) elif not fast_inference: @@ -2376,7 +2382,7 @@ class FastLlamaModel: token = token, max_position_embeddings = max_position_embeddings, trust_remote_code = trust_remote_code, - attn_implementation = "eager", + attn_implementation = preferred_attn_impl, **kwargs, ) model.fast_generate = make_fast_generate_wrapper(model.generate) @@ -2653,6 +2659,7 @@ class FastLlamaModel: loftq_config = {}, temporary_location = "_unsloth_temporary_saved_buffers", qat_scheme = None, + target_parameters = None, # For MoE expert layers (nn.Parameter) ensure_weight_tying = False, **kwargs, ): @@ -2683,6 +2690,7 @@ class FastLlamaModel: init_lora_weights = init_lora_weights, loftq_config = loftq_config, temporary_location = temporary_location, + target_parameters = target_parameters, ensure_weight_tying = ensure_weight_tying, **kwargs, ) @@ -2693,10 +2701,12 @@ class FastLlamaModel: return model transformers_set_seed(random_state) - if use_gradient_checkpointing == "unsloth": - patch_unsloth_smart_gradient_checkpointing( - dtype = model.get_input_embeddings().weight.dtype - ) + # Apply gradient checkpointing with smart heuristics + max_seq = getattr(model, "max_seq_length", 512) + dtype = model.get_input_embeddings().weight.dtype + use_gradient_checkpointing = apply_unsloth_gradient_checkpointing( + use_gradient_checkpointing, max_seq, dtype + ) if type(r) is not int: raise TypeError(f"Unsloth: Rank of {str(r)} must be an integer.") @@ -2966,6 +2976,10 @@ class FastLlamaModel: # Does not get lora yet, so get name from model, not base model is_classification = "Classification" in str(type(model)) + # Auto-detect MoE models and populate target_parameters for expert layers + if target_parameters is None: + target_parameters = get_moe_target_parameters(model, target_modules) + arguments = dict( r = r, lora_alpha = lora_alpha, @@ -2978,6 +2992,7 @@ class FastLlamaModel: loftq_config = loftq_config, use_rslora = use_rslora, modules_to_save = modules_to_save, + target_parameters = target_parameters, ensure_weight_tying = ensure_weight_tying, **kwargs, ) @@ -3081,35 +3096,17 @@ class FastLlamaModel: print("Unsloth: Training embed_tokens in mixed precision to save VRAM") assert hasattr(model.get_input_embeddings(), "modules_to_save") - new_dtype = ( - model.get_input_embeddings().modules_to_save.default.weight.dtype + _offload_frozen_module_for_training( + model.get_input_embeddings(), DEVICE_TYPE_TORCH, offload_device = None ) - if new_dtype == torch.float16: - # See https://github.com/unslothai/unsloth/pull/1200 - # Tesla T4 must use float32 and not float16 - new_dtype = torch.float32 - - model.get_input_embeddings().modules_to_save.default.to( - device = DEVICE_TYPE_TORCH, dtype = new_dtype, non_blocking = True - ) - model.get_input_embeddings().modules_to_save.default.requires_grad_(True) if train_lm_head: print("Unsloth: Training lm_head in mixed precision to save VRAM") assert hasattr(model.get_output_embeddings(), "modules_to_save") - new_dtype = ( - model.get_output_embeddings().modules_to_save.default.weight.dtype + _offload_frozen_module_for_training( + model.get_output_embeddings(), DEVICE_TYPE_TORCH, offload_device = None ) - if new_dtype == torch.float16: - # See https://github.com/unslothai/unsloth/pull/1200 - # Tesla T4 must use float32 and not float16 - new_dtype = torch.float32 - - model.get_output_embeddings().modules_to_save.default.to( - device = DEVICE_TYPE_TORCH, dtype = new_dtype, non_blocking = True - ) - model.get_output_embeddings().modules_to_save.default.requires_grad_(True) # Patch tokenizer to pad to the right internal_model = model diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index eb3b21e206..fd869c7b5f 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -88,7 +88,7 @@ from ._utils import ( patch_compiling_bitsandbytes, patch_model_and_tokenizer, prepare_model_for_kbit_training, - patch_unsloth_smart_gradient_checkpointing, + apply_unsloth_gradient_checkpointing, patch_compiled_autograd, process_vision_info, unsloth_compile_transformers, @@ -559,8 +559,10 @@ class FastLanguageModel(FastLlamaModel): **kwargs, ) - if use_gradient_checkpointing == "unsloth": - patch_unsloth_smart_gradient_checkpointing(dtype = dtype) + # Apply gradient checkpointing with smart heuristics + use_gradient_checkpointing = apply_unsloth_gradient_checkpointing( + use_gradient_checkpointing, max_seq_length, dtype + ) # Check if this is local model since the tokenizer gets overwritten if ( @@ -734,6 +736,7 @@ class FastModel(FastBaseModel): qat_scheme = None, load_in_fp8 = False, # fp8 LoRA (True, False, 'block') unsloth_tiled_mlp = False, + target_parameters = None, # For MoE expert parameters *args, **kwargs, ): @@ -1188,9 +1191,10 @@ class FastModel(FastBaseModel): os.environ["UNSLOTH_FORCE_FLOAT32"] = "1" dtype = torch.bfloat16 # Change to bfloat16 loading break - # Patch gradient checkpointing - if use_gradient_checkpointing == "unsloth": - patch_unsloth_smart_gradient_checkpointing(dtype = dtype) + # Apply gradient checkpointing with smart heuristics + use_gradient_checkpointing = apply_unsloth_gradient_checkpointing( + use_gradient_checkpointing, max_seq_length, dtype + ) with redirector: patch_loss_functions(torch_compile = False) model_types, supports_sdpa = unsloth_compile_transformers( diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py old mode 100644 new mode 100755 index 9788207c99..eacfecc6c3 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -26,6 +26,7 @@ from unsloth_zoo.compiler import create_new_function from unsloth_zoo.log import logger from unsloth_zoo.logging_utils import PatchRLStatistics from unsloth_zoo.rl_replacements import RL_REPLACEMENTS +from ..device_type import DEVICE_TYPE from .rl_replacements import ( RL_EXTRA_ARGS, RL_FUNCTIONS, @@ -251,6 +252,7 @@ from torch.nn import functional as F import inspect from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling as TransformersDataCollatorForLanguageModeling from transformers.training_args import ParallelMode +from unsloth_zoo.device_type import DEVICE_TYPE, device_synchronize # Wrap trainer with padding to right and enable training mode # Also patches W&B since multiple runs must use wandb.finish() @@ -355,6 +357,7 @@ class Unsloth{RLConfig_name}({RLConfig_name}): ) self.unsloth_logit_chunk_multiplier = unsloth_logit_chunk_multiplier {max_seq_length_post} +{RLConfig_post} pass {RLTrainer_extras} @@ -417,7 +420,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): try: trainer = eval(f"trl.trainer.{trainer_file}") except Exception as error: - print(f"Unsloth: Could not import trl.trainer.{trainer_file}: {error}") + logger.info(f"Unsloth: Could not import trl.trainer.{trainer_file}: {error}") return # Get SFTTrainer and SFTConfig names @@ -888,6 +891,15 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): ) extra_args += learning_rate_check + # Fix num_train_epochs = None causing TypeError in Trainer.__init__ + # Trainer does `args.num_train_epochs > 0` which fails when None + if "num_train_epochs" in call_args: + num_train_epochs_check = ( + "if num_train_epochs is None:\n" + " num_train_epochs = 3.0 # Default to 3 epochs if None, max_steps will override\n" + ) + extra_args += num_train_epochs_check + # Check if max_seq_length is NOT defined (max_length is now default) if "max_seq_length" not in call_args and "max_length" in call_args: max_seq_length_pre = """max_seq_length : Optional[int] = field( @@ -1023,6 +1035,18 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): RLConfig_extra_args = extra_args RLConfig_call_args = call_args + # TRL 0.27.0+ forces use_reentrant=False in gradient_checkpointing_kwargs. + # Unsloth gradient checkpointing requires use_reentrant=True, so we remove + # the setting after super().__init__() when it gets auto-applied. + RLConfig_post = "" + if trl_version >= Version("0.27.0") and RLConfig_name == "GRPOConfig": + RLConfig_post = ( + " # Unsloth: Remove use_reentrant=False forced by TRL 0.27.0+\n" + " if getattr(self, 'gradient_checkpointing_kwargs', None) is not None:\n" + " if 'use_reentrant' in self.gradient_checkpointing_kwargs:\n" + " del self.gradient_checkpointing_kwargs['use_reentrant']\n" + ) + # Patch vLLM and other functions RLTrainer_extras = patch_functions( RLTrainer, trainer_file, RLTrainer_name, all_imports, imports @@ -1075,6 +1099,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): RLConfig_extra_args = RLConfig_extra_args, RLConfig_call_args = RLConfig_call_args, RLConfig_kwargs = ",**kwargs"[1 if RLConfig_call_args.endswith(",") else 0 :], + RLConfig_post = RLConfig_post, RLTrainer_extras = RLTrainer_extras, RLTrainer_post = RLTrainer_post, RL_pre = RL_pre, @@ -1090,6 +1115,68 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): align_logprobs_with_mask_code = align_logprobs_with_mask_code, ) + if RLTrainer_name == "GRPOTrainer": + # Base torch_compile_options shared by all device types + base_options = """torch_compile_options = { + "epilogue_fusion" : True, + "max_autotune" : False, + "shape_padding" : True, + "trace.enabled" : False,""" + + # Generate torch_compile_options based on device type + if DEVICE_TYPE == "cuda": + # CUDA-specific options (added to base options) + new_options = ( + base_options + + """ + "triton.enable_persistent_tma_matmul": torch.cuda.get_device_capability()[0] >= 9, + "cuda.cutlass_epilogue_fusion_enabled": torch.cuda.get_device_capability()[0] >= 9, + "cuda.cutlass_tma_only": torch.cuda.get_device_capability()[0] >= 9, + "cuda.compile_opt_level" : "-O2", + "cuda.enable_cuda_lto" : True, + }""" + ) + else: + # XPU, HIP, and other device types use base options only + new_options = ( + base_options + + """ + }""" + ) + + pattern = r"torch_compile_options\s*=\s*\{[^}]*\}" + + RLTrainer_source = re.sub( + pattern, new_options, RLTrainer_source, flags = re.DOTALL + ) + + if trl_version >= Version("0.27.0"): + peft_pattern = ( + r"\s*if is_peft_available\(\) and is_peft_model\(model\) and args\.beta != 0\.0:" + r".*?" + r"param\.data = param\.data\.to\(torch\.bfloat16\)" + ) + + replacement_comment = "\n # PEFT initialization logic removed via script for trl >= 0.27.0\n" + + RLTrainer_source = re.sub( + peft_pattern, replacement_comment, RLTrainer_source, flags = re.DOTALL + ) + + elif trl_version >= Version("0.26.0"): + peft_block_pattern = ( + r"\s*if is_peft_available\(\) and isinstance\(model, PeftModel\) and peft_config is not None:" + r".*?" + r"param\.data = param\.data\.to\(torch\.bfloat16\)" + ) + + RLTrainer_source = re.sub( + peft_block_pattern, + "\n # TRL PEFT 0.26.0 initialization logic removed on unsloth side.\n", + RLTrainer_source, + flags = re.DOTALL, + ) + if RLTrainer_name == "SFTTrainer": original_text = 'self._signature_columns = ["input_ids", "attention_mask", "completion_mask"]' new_text = 'self._signature_columns = ["input_ids", "attention_mask", "completion_mask","labels"]' @@ -1193,6 +1280,8 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import init = init.replace( "model = self._prepare_peft_model(model, peft_config, args)\n", "pass\n" ) + # TRL 0.22.0+ uses prepare_peft_model as a standalone function + init = init.replace("model = prepare_peft_model(model, peft_config, args)", "pass") # Skip add_adapter("ref") for reference model computation # Unsloth: We comment out the "ref" adapter creation because: diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py old mode 100644 new mode 100755 index ff36da125d..8208dc922a --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -27,8 +27,10 @@ import inspect from collections import defaultdict from unsloth_zoo.rl_replacements import RL_REPLACEMENTS, left_pack_padding from unsloth_zoo.utils import Version +from trl import __version__ as trl_version_raw from importlib.metadata import version as importlib_version from unsloth_zoo.log import logger +from unsloth_zoo.device_type import device_synchronize import importlib.util from ..device_type import ( is_hip, @@ -56,6 +58,14 @@ torch_compile_options = { "triton.cudagraphs": False, } +try: + trl_version = Version(trl_version_raw) +except Exception: + try: + trl_version = Version(importlib_version("trl")) + except Exception: + trl_version = Version("0.0.0") + # Check untrained tokens def sft_trainer_fix_untrained_tokens(call_args, extra_args): @@ -75,6 +85,16 @@ def sft_trainer_fix_untrained_tokens(call_args, extra_args): RL_EXTRA_ARGS["sft_trainer"].append(sft_trainer_fix_untrained_tokens) +# Fix top_k for GRPO vLLM. +# https://github.com/huggingface/trl/pull/4695 with this change trl added top_k in GRPOConfig and defaults to 0 +# We don't want that since vllm's all include top_k is -1 and 0 returns an error on SamplingParams creation. +def grpo_config_fix_vllm_top_k(old_RLTrainer_source, old_RLConfig_source): + return "if use_vllm and (top_k is None or top_k == 0): top_k = -1\n" + + +RL_CONFIG_CHANGES["grpo_trainer"].append(grpo_config_fix_vllm_top_k) + + # Remove DPO columns which might randomnly be tokenized def dpo_trainer_fix_columns(call_args, extra_args): if "model" in call_args and "train_dataset" in call_args: @@ -236,6 +256,30 @@ def grpo_trainer__generate_single_turn(function_name, function): "", function, ) + + # TRL 0.24.0-0.25.1 truncation regression fix + # + # TRL 0.22.2-0.23.1 used smart truncation via truncate_with_protected_tokens(): + # - Tokenizes first without truncation + # - Then truncates keeping the RIGHTMOST tokens (preserves assistant turn) + # - Protects special tokens (image_token, vision_start/end) from removal + # + # TRL 0.24.0-0.25.1 removed this and passed kwargs directly to the tokenizer: + # max_length=self.max_prompt_length, truncation=True, add_special_tokens=False + # This causes issues because tokenizer truncation doesn't protect special tokens + # and may not preserve the end of the prompt properly. + # + # TRL 0.26.2+ removed these kwargs entirely (no tokenizer-level truncation). + # + # Fix: Remove these kwargs so TRL 0.24.0-0.25.1 behaves like 0.26.2+ (no truncation). + # This is a no-op for versions that don't have these kwargs (0.22.2-0.23.1, 0.26.2+). + for pattern in [ + r'["\']?max_length["\']?\s*[:=]\s*self\.max_prompt_length\s*,\s*\n?', + r'["\']?truncation["\']?\s*[:=]\s*True\s*,\s*\n?', + r'["\']?add_special_tokens["\']?\s*[:=]\s*False\s*,\s*\n?', + ]: + function = re.sub(pattern, "", function) + return function @@ -283,7 +327,7 @@ def grpo_trainer__generate_and_score_completions(function_name, function): re.MULTILINE, ) - replacement_text = """ + replacement_text = """ if self.args.gradient_accumulation_steps % generate_every != 0 or ( self.use_vllm ):""" @@ -365,7 +409,7 @@ def grpo_trainer__generate_and_score_completions(function_name, function): replacement_string = """ if "image_sizes" in prompt_inputs: output["image_sizes"] = prompt_inputs["image_sizes"] if max_left_pad is not None: - output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1) + output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1) try: if self.use_vllm and getattr(self, "vllm_importance_sampling_correction", False): output["sampling_per_token_logps"] = sampling_per_token_logps @@ -374,6 +418,31 @@ def grpo_trainer__generate_and_score_completions(function_name, function): function = function.replace(string_to_find, replacement_string) + # TRL 0.24.0+ extracts prompts = [x["prompt"] for x in inputs], losing metadata + # like reasoning_effort. Inject code to store per-sample chat_template_kwargs on self. + _metadata_extraction = ( + "\n" + " # Unsloth: Extract per-sample chat_template_kwargs before metadata is lost\n" + " _ct_ = getattr(self.processing_class, 'chat_template', None) or ''\n" + " _sk_ = {'prompt', 'chosen', 'rejected', 'completion', 'messages', 'label',\n" + " 'images', 'image', 'videos', 'video', 'audios', 'audio'}\n" + " self._unsloth_batch_chat_kwargs = []\n" + " for _inp_ in inputs:\n" + " _kw_ = {}\n" + " if isinstance(_inp_, dict):\n" + " for _k_ in _inp_.keys() - _sk_:\n" + " if _k_ in _ct_ and isinstance(_inp_[_k_], str):\n" + " _kw_[_k_] = _inp_[_k_]\n" + " self._unsloth_batch_chat_kwargs.append(_kw_)\n" + ) + # Insert after: prompts = [x["prompt"] for x in inputs] + _target_line = 'prompts = [x["prompt"] for x in inputs]' + if _target_line in function: + function = function.replace( + _target_line, + _target_line + _metadata_extraction, + ) + # This path is for TRL 0.24.0 images is a variable exclusive to this version string_to_find = """ if images is not None: output["num_images"] = num_images""" @@ -381,7 +450,7 @@ def grpo_trainer__generate_and_score_completions(function_name, function): replacement_string = """ if images is not None: output["num_images"] = num_images if max_left_pad is not None: - output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1) + output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1) try: if self.use_vllm and getattr(self, "vllm_importance_sampling_correction", False): output["sampling_per_token_logps"] = sampling_per_token_logps @@ -390,6 +459,17 @@ def grpo_trainer__generate_and_score_completions(function_name, function): function = function.replace(string_to_find, replacement_string) + if trl_version >= Version("0.25.0"): + # We replace the call using 'completions' with one using 'completions_text' + string_to_find = " rewards_per_func = self._calculate_rewards(inputs, prompts, completions, completion_ids_list)" + replacement_string = ( + " if images is not None:\n" + " rewards_per_func = self._calculate_rewards(inputs, prompts_text, completions_text, completion_ids_list)\n" + " else:\n" + " rewards_per_func = self._calculate_rewards(inputs, prompts, completions, completion_ids_list)" + ) + function = function.replace(string_to_find, replacement_string) + if "wake_up()" not in function: # Sleep functionality has been added to trl in v0.23.0. We do not want to redo this. # https://github.com/huggingface/trl/commit/edbe8234bc7e528f72ac76607de9d3e4753e2709 @@ -434,9 +514,10 @@ def grpo_trainer_fix_maybe_apply_chat_template(function_name, function): _chat_template_ = getattr(self.processing_class, "chat_template", None) if _chat_template_ is None: _chat_template_ = "" _supported_keys_ = set(("prompt", "chosen", "rejected", "completion", "messages", "label")) + _batch_chat_kwargs_ = getattr(self, "_unsloth_batch_chat_kwargs", None) prompts_text = [] - for _example_ in __INPUTS__REPLACEMENT__: + for _idx_, _example_ in enumerate(__INPUTS__REPLACEMENT__): _tokenizer_kwargs_ = {} if type(_example_) is not dict: _example_ = {"prompt": _example_} @@ -446,6 +527,10 @@ def grpo_trainer_fix_maybe_apply_chat_template(function_name, function): v = _example_[k] if type(v) is str: _tokenizer_kwargs_[k] = v + if _batch_chat_kwargs_ is not None and _idx_ < len(_batch_chat_kwargs_): + for _bk_, _bv_ in _batch_chat_kwargs_[_idx_].items(): + if _bk_ not in _tokenizer_kwargs_: + _tokenizer_kwargs_[_bk_] = _bv_ _x_ = maybe_apply_chat_template(_example_, self.processing_class, **_tokenizer_kwargs_)["prompt"] prompts_text.append(_x_) """ @@ -771,7 +856,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): ) # This is needed to avoid race conditions with GPT OSS offload_embbed=True # However, it seems that this line does not slow down or disrupt models. - torch.cuda.synchronize() + device_synchronize() all_logprobs_list.append(logprobs_chunk) logprobs = torch.cat(all_logprobs_list, dim = 0) entropies = None @@ -914,7 +999,7 @@ def grpo_trainer_compute_loss(function_name, function): max_left_pad = inputs.get("max_left_pad", 0) if per_token_logps is not None: - loss, completion_length, mean_kl, delta, flat_is_ratio = ( + loss, completion_length, mean_kl, delta, flat_is_ratio, coef_1 = ( grpo_compute_loss_slow( ref_logps, per_token_logps, @@ -944,7 +1029,7 @@ def grpo_trainer_compute_loss(function_name, function): ) else: if hasattr(self.args, "loss_type"): - loss, completion_length, mean_kl, delta, flat_is_ratio = ( + loss, completion_length, mean_kl, delta, flat_is_ratio, coef_1 = ( grpo_accumulated_loss( trainer = self, input_ids = _input_ids, @@ -976,7 +1061,7 @@ def grpo_trainer_compute_loss(function_name, function): ) else: # to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17 - loss, completion_length, mean_kl = grpo_accumulated_loss( + loss, completion_length, mean_kl, coef_1 = grpo_accumulated_loss( trainer = self, input_ids = _input_ids, logits_to_keep = logits_to_keep, @@ -991,7 +1076,6 @@ def grpo_trainer_compute_loss(function_name, function): logit_scale_divide = logit_scale_divide, attention_mask = attention_mask, ) - if "train" in self._metrics: mode = "eval" if self.control.should_evaluate else "train" self._metrics[mode]["completion_length"].append(completion_length.item()) @@ -1053,6 +1137,53 @@ def grpo_trainer_compute_loss(function_name, function): .item() ) + completion_token_count = completion_mask.sum().clamp(min = 1.0) + + def masked_batch_mean(x): + if x.shape[1] == 1: # when importance_sampling_level == "sequence" + return x.mean() + else: + return (x * completion_mask).sum() / completion_token_count + + if advantages.dim() == 1: + advantages = advantages.unsqueeze(1) + + if self.loss_type in ["grpo", "bnpo", "dr_grpo", "dapo"]: + # Compute the clipped probability ratios + is_low_clipped = (coef_1 < 1 - self.epsilon_low) & (advantages < 0) + is_high_clipped = (coef_1 > 1 + self.epsilon_high) & (advantages > 0) + is_region_clipped = is_low_clipped | is_high_clipped + + low_clip = masked_batch_mean(is_low_clipped.float()) + high_clip = masked_batch_mean(is_high_clipped.float()) + clip_ratio = masked_batch_mean(is_region_clipped.float()) + + gathered_low_clip = self.accelerator.gather(low_clip) + self._metrics[mode]["clip_ratio/low_mean"].append( + gathered_low_clip.nanmean().item() + ) + self._metrics[mode]["clip_ratio/low_min"].append( + nanmin(gathered_low_clip).item() + ) + gathered_high_clip = self.accelerator.gather(high_clip) + self._metrics[mode]["clip_ratio/high_mean"].append( + gathered_high_clip.nanmean().item() + ) + self._metrics[mode]["clip_ratio/high_max"].append( + nanmax(gathered_high_clip).item() + ) + gathered_clip_ratio = self.accelerator.gather(clip_ratio) + self._metrics[mode]["clip_ratio/region_mean"].append( + gathered_clip_ratio.nanmean().item() + ) + elif self.loss_type == "cispo": + is_cispo_clipped = (coef_1 > self.epsilon_high) & (advantages > 0) + cispo_clip_ratio = masked_batch_mean(is_cispo_clipped.float()) + gathered_cispo_clip_ratio = self.accelerator.gather(cispo_clip_ratio) + self._metrics[mode]["cispo_clip_ratio"].append( + gathered_cispo_clip_ratio.nanmean().item() + ) + return loss function = inspect.getsource(compute_loss) diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 6835f2e986..93811c2668 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -98,6 +98,7 @@ VLLM_SUPPORTED_VLM = [ "gemma3", "mistral3", "qwen3_vl", + "qwen3_vl_moe", ] VLLM_NON_LORA_VLM = [ "mllama", @@ -517,9 +518,23 @@ class FastBaseModel: correct_dtype = None # Stop SDPA for some archs like Pixtral / Mistral3 + flex_attn_impl = None + if auto_config is None: + auto_config = AutoConfig.from_pretrained( + model_name, + token = token, + trust_remote_code = trust_remote_code, + ) + try: + model_class = auto_model._model_mapping[auto_config.__class__] + except Exception: + model_class = None + flex_attn_impl = prefer_flex_attn_if_supported(model_class, auto_config) + + default_attn_impl = "flex_attention" if flex_attn_impl else "sdpa" if not ("attn_implementation" in kwargs): - kwargs["attn_implementation"] = "sdpa" - if not supports_sdpa: + kwargs["attn_implementation"] = default_attn_impl + if not supports_sdpa and kwargs.get("attn_implementation") == "sdpa": if os.environ.get("UNSLOTH_ENABLE_FLEX_ATTENTION", "0") == "0": print( f"Unsloth: {model_type_arch.title()} does not support SDPA - switching to fast eager." @@ -651,12 +666,19 @@ class FastBaseModel: kwargs = add_dtype_kwargs(torch_dtype, kwargs) - model_config = AutoConfig.from_pretrained( - model_name, - token = token, - attn_implementation = "sdpa" if supports_sdpa else "eager", - trust_remote_code = trust_remote_code, - ) + config_attn_impl = kwargs.get("attn_implementation", None) + if config_attn_impl is None: + config_attn_impl = "sdpa" if supports_sdpa else "eager" + if auto_config is None: + auto_config = AutoConfig.from_pretrained( + model_name, + token = token, + trust_remote_code = trust_remote_code, + ) + setattr(auto_config, "_attn_implementation", config_attn_impl) + if hasattr(auto_config, "attn_implementation"): + setattr(auto_config, "attn_implementation", config_attn_impl) + model_config = auto_config verify_fp8_support_if_applicable(model_config) raise_handler = RaiseUninitialized() @@ -939,6 +961,7 @@ class FastBaseModel: task_type = TaskType.CAUSAL_LM, temporary_location = "_unsloth_temporary_saved_buffers", qat_scheme = None, + target_parameters = None, # For MoE expert layers (nn.Parameter) ensure_weight_tying = False, # [TODO] Add `ensure_weight_tying` for `modules_to_save` for vision models **kwargs, ): @@ -1020,6 +1043,10 @@ class FastBaseModel: loftq_config, lora_dropout, bias, init_lora_weights, model ) + # Auto-detect MoE models and populate target_parameters for expert layers + if target_parameters is None: + target_parameters = get_moe_target_parameters(model, target_modules) + # Get only allowed parameters for LoraConfig local_variables = { **locals(), diff --git a/unsloth/tokenizer_utils.py b/unsloth/tokenizer_utils.py index 0136e3498e..087a9a7f8a 100644 --- a/unsloth/tokenizer_utils.py +++ b/unsloth/tokenizer_utils.py @@ -602,7 +602,10 @@ def load_correct_tokenizer( old_chat_template = getattr(tokenizer, "chat_template", None) # Ignore mistral type models since they don't have an add_generation_prompt - if "mistral" in str(getattr(tokenizer, "name_or_path", "")).lower(): + if any( + s in str(getattr(tokenizer, "name_or_path", "")).lower() + for s in ["mistral", "qwen3guard"] + ): chat_template = old_chat_template # Also check Llama-2 old style models