Merge branch 'main' into pip

This commit is contained in:
Daniel Han 2026-02-05 06:10:26 -08:00
commit 96f14c999a
23 changed files with 2006 additions and 201 deletions

View file

@ -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.

View file

@ -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:

View file

@ -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.
<details>
<summary>Click for more news</summary>
- **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
</details>
### 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')
```

View file

@ -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":

View file

@ -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')

View file

@ -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

View file

@ -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 <https://www.gnu.org/licenses/>.
"""
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

View file

@ -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

View file

@ -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 <https://www.gnu.org/licenses/>.
"""
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():

View file

@ -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)

View file

@ -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",

View file

@ -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,

View file

@ -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,

View file

@ -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,

View file

@ -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,

View file

@ -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

450
unsloth/models/glm4_moe.py Normal file
View file

@ -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,
)

View file

@ -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

View file

@ -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(

91
unsloth/models/rl.py Normal file → Executable file
View file

@ -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:

149
unsloth/models/rl_replacements.py Normal file → Executable file
View file

@ -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)

View file

@ -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(),

View file

@ -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