diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS
new file mode 100644
index 0000000000..f4d18645f2
--- /dev/null
+++ b/.github/CODEOWNERS
@@ -0,0 +1,46 @@
+# Inspired from https://github.com/vllm-project/vllm/blob/main/.github/CODEOWNERS
+
+/unsloth/models/loader.py @danielhanchen @mmathew23
+/unsloth/models/llama.py @Datta0 @danielhanchen @mmathew23
+/unsloth/models/rl.py @Datta0 @pluesclues @danielhanchen
+/unsloth/models/rl_replacements.py @Datta0 @pluesclues @danielhanchen
+/unsloth/trainer.py @danielhanchen
+/unsloth/models/sentence_transformer.py @Etherll @danielhanchen
+/unsloth/save.py @rolandtannous @danielhanchen
+/unsloth/tokenizer_utils.py @mmathew23 @danielhanchen
+/unsloth/chat_templates.py @rolandtannous @danielhanchen
+/unsloth/ollama_template_mappers.py @rolandtannous @danielhanchen
+/unsloth/kernels/moe/*.py @Datta0
+/unsloth/import_fixes.py @danielhanchen
+/unsloth/device_type.py @danielhanchen
+/unsloth/_auto_install.py @danielhanchen
+/unsloth/dataprep/*.py @danielhanchen
+/unsloth/kernels/cross_entropy_loss.py @danielhanchen
+/unsloth/kernels/fast_lora.py @danielhanchen
+/unsloth/kernels/flex_attention.py @danielhanchen
+/unsloth/kernels/fp8.py @Datta0
+/unsloth/kernels/geglu.py @danielhanchen
+/unsloth/kernels/layernorm.py @danielhanchen
+/unsloth/kernels/rms_layernorm.py @danielhanchen
+/unsloth/kernels/rope_embedding.py @danielhanchen
+/unsloth/kernels/swiglu.py @danielhanchen
+/unsloth/kernels/utils.py @danielhanchen @Datta0
+/unsloth/models/_utils.py @danielhanchen @mmathew23
+/unsloth/models/cohere.py @danielhanchen
+/unsloth/models/dpo.py @danielhanchen
+/unsloth/models/falcon_h1.py @danielhanchen
+/unsloth/models/gemma.py @danielhanchen
+/unsloth/models/gemma2.py @danielhanchen
+/unsloth/models/glm4_moe.py @Datta0
+/unsloth/models/granite.py @danielhanchen
+/unsloth/models/llama4.py @danielhanchen
+/unsloth/models/loader_utils.py @Datta0 @danielhanchen
+/unsloth/models/mapper.py @danielhanchen
+/unsloth/models/mistral.py @danielhanchen
+/unsloth/models/qwen2.py @danielhanchen
+/unsloth/models/qwen3.py @Datta0
+/unsloth/models/qwen3_moe.py @Datta0
+/unsloth/models/vision.py @mmathew23 @danielhanchen
+/unsloth/utils/attention_dispatch.py @mmathew23
+/unsloth/utils/hf_hub.py @mmathew23
+/unsloth/utils/packing.py @mmathew23
diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml
index 47c2dd8010..5c822d9b38 100644
--- a/.pre-commit-config.yaml
+++ b/.pre-commit-config.yaml
@@ -1,6 +1,6 @@
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
- rev: v0.15.0
+ rev: v0.15.4
hooks:
- id: ruff
args:
diff --git a/README.md b/README.md
index bfe7f2255b..1314cb1c59 100644
--- a/README.md
+++ b/README.md
@@ -22,12 +22,11 @@ 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 |
+| **Qwen3.5 (4B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_5_(4B)_Vision.ipynb) | 1.5x faster | 60% less |
+| **gpt-oss (20B)** | [▶️ Start for free](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/gpt-oss-(20B)-Fine-tuning.ipynb) | 2x faster | 70% 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 (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 |
@@ -49,19 +48,20 @@ For Windows, `pip install unsloth` works only if you have Pytorch installed. Rea
### Docker
Use our official [Unsloth Docker image](https://hub.docker.com/r/unsloth/unsloth) ```unsloth/unsloth``` container. Read our [Docker Guide](https://unsloth.ai/docs/get-started/install/docker).
-### Blackwell & DGX Spark
-For RTX 50x, B200, 6000 GPUs: `pip install unsloth`. Read our [Blackwell Guide](https://unsloth.ai/docs/blog/fine-tuning-llms-with-blackwell-rtx-50-series-and-unsloth) and [DGX Spark Guide](https://unsloth.ai/docs/blog/fine-tuning-llms-with-nvidia-dgx-spark-and-unsloth) for more details.
+### AMD, Intel, Blackwell & DGX Spark
+For RTX 50x, B200, 6000 GPUs: `pip install unsloth`. Read our guides for: [Blackwell](https://unsloth.ai/docs/blog/fine-tuning-llms-with-blackwell-rtx-50-series-and-unsloth) and [DGX Spark](https://unsloth.ai/docs/blog/fine-tuning-llms-with-nvidia-dgx-spark-and-unsloth).
+To install Unsloth on **AMD** and **Intel** GPUs, follow our [AMD Guide](https://unsloth.ai/docs/get-started/install/amd) and [Intel Guide](https://unsloth.ai/docs/get-started/install/intel).
## 🦥 Unsloth News
+- **Qwen3.5** - 0.8B, 2B, 4B, 9B, 27B, 35-A3B, 112B-A10B are now supported. [Guide + notebooks](https://unsloth.ai/docs/models/qwen3.5/fine-tune)
+- Train **MoE LLMs 12x faster** with 35% less VRAM - DeepSeek, GLM, Qwen and gpt-oss. [Blog](https://unsloth.ai/docs/new/faster-moe)
- **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)
- **500K Context**: Training a 20B model with >500K context is now possible on an 80GB GPU. [Blog](https://unsloth.ai/docs/blog/500k-context-length-fine-tuning)
-- **FP8 Reinforcement Learning**: You can now do FP8 GRPO on consumer GPUs. [Blog](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/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/tutorials/deepseek-ocr-how-to-run-and-fine-tune) • [Notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Deepseek_OCR_(3B).ipynb)
+- **FP8 & Vision RL**: You can now do FP8 & VLM GRPO on consumer GPUs. [FP8 Blog](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/fp8-reinforcement-learning) • [Vision RL](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/vision-reinforcement-learning-vlm-rl)
- **Docker**: Use Unsloth with no setup & environment issues with our new image. [Guide](https://unsloth.ai/docs/blog/how-to-fine-tune-llms-with-unsloth-and-docker) • [Docker image](https://hub.docker.com/r/unsloth/unsloth)
-- **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 [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.
+- **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 [Guide](https://unsloth.ai/docs/models/gpt-oss-how-to-run-and-fine-tune).
Click for more news
@@ -102,9 +102,9 @@ For RTX 50x, B200, 6000 GPUs: `pip install unsloth`. Read our [Blackwell Guide](
* 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.
-* Supports NVIDIA (since 2018), [AMD](https://unsloth.ai/docs/get-started/install/amd) and Intel GPUs. Minimum CUDA Capability 7.0 (V100, T4, Titan V, RTX 20, 30, 40x, A100, H100, L40 etc)
-* Works on **Linux**, WSL and **Windows**
+* Export and [deploy your model](https://unsloth.ai/docs/basics/inference-and-deployment) to [GGUF](https://unsloth.ai/docs/basics/inference-and-deployment/saving-to-gguf) llama.cpp, [vLLM](https://unsloth.ai/docs/basics/inference-and-deployment/vllm-guide), [SGLang](https://unsloth.ai/docs/basics/inference-and-deployment/sglang-guide) and Hugging Face.
+* Supports NVIDIA (since 2018), [AMD](https://unsloth.ai/docs/get-started/install/amd) and [Intel](https://unsloth.ai/docs/get-started/install/intel) GPUs. Minimum CUDA Capability 7.0 (V100, T4, Titan V, RTX 20, 30, 40x, A100, H100, L40 etc)
+* Works on **Linux**, WSL and **[Windows](https://unsloth.ai/docs/get-started/install/windows-installation)**
* All kernels written in OpenAI's Triton language. Manual backprop engine.
* If you trained a model with 🦥Unsloth, you can use this cool sticker!
@@ -124,25 +124,33 @@ pip install --upgrade --force-reinstall --no-cache-dir unsloth unsloth_zoo
```
See [here](#advanced-pip-installation) for advanced pip install instructions.
### Windows Installation
+For this method, we will be utilizing Anaconda. You can view the [full guide with screenshots here](https://unsloth.ai/docs/get-started/install/windows-installation).
+1. **Install Miniconda (or Anaconda):** Miniconda is recommended. Install [Miniconda](https://www.anaconda.com/docs/getting-started/miniconda/install) or [Anaconda](https://www.anaconda.com/download), then open Anaconda PowerShell Prompt to continue.
-1. **Install NVIDIA Video Driver:**
- You should install the latest driver for your GPU. Download drivers here: [NVIDIA GPU Driver](https://www.nvidia.com/Download/index.aspx).
+2. **Create a Conda Environment:** Create and activate a fresh Python 3.12 environment for Unsloth.
-2. **Install Visual Studio C++:**
- You will need Visual Studio, with C++ installed. By default, C++ is not installed with [Visual Studio](https://visualstudio.microsoft.com/vs/community/), so make sure you select all of the C++ options. Also select options for Windows 10/11 SDK. For detailed instructions with options, see [here](https://unsloth.ai/docs/get-started/install/windows-installation#method-3-windows-directly).
+ ```bash
+ conda create --name unsloth_env python==3.12 -y
+ conda activate unsloth_env
+ ```
-3. **Install CUDA Toolkit:**
- Follow the instructions to install [CUDA Toolkit](https://developer.nvidia.com/cuda-toolkit-archive).
+3. **Check Your GPU and CUDA Version:** Run `nvidia-smi` to confirm that your NVIDIA GPU is detected and note the CUDA version shown in the output. If `nvidia-smi` does not work, reinstall the latest [NVIDIA drivers](https://www.nvidia.com/en-us/drivers/).
-4. **Install PyTorch:**
- You will need the correct version of PyTorch that is compatible with your CUDA drivers, so make sure to select them carefully.
- [Install PyTorch](https://pytorch.org/get-started/locally/).
+4. **Install PyTorch:** Install the Windows pip build of PyTorch that matches your CUDA version. Use [Install PyTorch](https://pytorch.org/get-started/locally/) to select the correct command for your system, then verify that PyTorch can see your GPU.
-5. **Install Unsloth:**
-
-```python
-pip install unsloth
-```
+ ```python
+ import torch
+ print(torch.cuda.is_available())
+ A = torch.ones((10, 10), device="cuda")
+ B = torch.ones((10, 10), device="cuda")
+ A @ B
+ ```
+
+5. **Install Unsloth:** Only install Unsloth after PyTorch is working correctly.
+
+ ```bash
+ pip install unsloth
+ ```
#### Advanced/Troubleshooting
For **advanced installation instructions** or if you see weird errors during installations:
diff --git a/pyproject.toml b/pyproject.toml
index f67f9ad121..9d410b4033 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -7,7 +7,7 @@ name = "unsloth"
dynamic = ["version"]
description = "2-5X faster training, reinforcement learning & finetuning"
readme = "README.md"
-requires-python = ">=3.9,<3.14"
+requires-python = ">=3.9,<3.15"
license = "Apache-2.0"
keywords = ["ai", "llm", "reinforcement learning", "machine learning", "artificial intelligence", "pytorch"]
authors = [
@@ -55,13 +55,13 @@ huggingfacenotorch = [
"huggingface_hub>=0.34.0",
"hf_transfer",
"diffusers",
- "transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,!=4.57.4,!=4.57.5,<=4.57.6",
+ "transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,!=4.57.4,!=4.57.5,!=5.0.0,!=5.1.0,<=5.2.0",
"trl>=0.18.2,!=0.19.0,<=0.24.0",
"sentence-transformers",
]
huggingface = [
"unsloth[huggingfacenotorch]",
- "unsloth_zoo>=2026.2.1",
+ "unsloth_zoo>=2026.3.1",
"torchvision",
"unsloth[triton]",
]
@@ -551,10 +551,10 @@ colab-ampere-torch220 = [
"flash-attn>=2.6.3 ; ('linux' in sys_platform)",
]
colab-new = [
- "unsloth_zoo>=2026.2.1",
+ "unsloth_zoo>=2026.3.1",
"packaging",
"tyro",
- "transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,!=4.57.4,!=4.57.5,<=4.57.6",
+ "transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,!=4.57.4,!=4.57.5,!=5.0.0,!=5.1.0,<=5.2.0",
"datasets>=3.4.1,!=4.0.*,!=4.1.0,<4.4.0",
"sentencepiece>=0.2.0",
"tqdm",
@@ -989,14 +989,50 @@ intelgputorch290 = [
intel-gpu-torch290 = [
"unsloth[intelgputorch290]"
]
+intelgputorch210 = [
+ "unsloth_zoo[intelgpu]",
+ "unsloth[huggingfacenotorch]",
+
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=c169a1de14c19673b17c751290d467fa282fc90fa5da4314b2e5cdab1f553146 ; platform_system == 'Linux' and python_version == '3.10' and platform_machine == 'x86_64'",
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=013d9dd5d6479bd22983161f462e61c8dbe1d82e6730624a7a8d5945507eaa61 ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=afc8cabfbf7ed51fd278d1e0f88d6afc157b0201bad4b99d681e4d542f9e66d4 ; platform_system == 'Linux' and python_version == '3.12' and platform_machine == 'x86_64'",
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=0d24c1716088f2764d0d24c64227732195b6a42706c3c5fc89eeb4904bfa0818 ; platform_system == 'Linux' and python_version == '3.13' and platform_machine == 'x86_64'",
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp310-cp310-win_amd64.whl#sha256=c83ab007311d9cfb6e809ee5a4587d99a9eef4be720b90da4f1aaa68b45139a0 ; sys_platform == 'win32' and python_version == '3.10' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp311-cp311-win_amd64.whl#sha256=debf75348da8e8c7166b4d4a9b91d1508bb8d6581e339f79f7604b2e6746bacd ; sys_platform == 'win32' and python_version == '3.11' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp312-cp312-win_amd64.whl#sha256=97337a47425f1963a723475bd61037460e84ba01db4f87a1d662c3718ff6c47e ; sys_platform == 'win32' and python_version == '3.12' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "pytorch_triton_xpu @ https://download.pytorch.org/whl/pytorch_triton_xpu-3.5.0-cp313-cp313-win_amd64.whl#sha256=2caf8138695f6abb023ecd02031a2611ba1bf8fff2f19802567cb2fadefe9e87 ; sys_platform == 'win32' and python_version == '3.13' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp310-cp310-linux_x86_64.whl#sha256=abb1d1ec1ac672bac0ff35420c965f2df0c636ef9d94e2a830e34578489d0a57 ; platform_system == 'Linux' and python_version == '3.10' and platform_machine == 'x86_64'",
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp311-cp311-linux_x86_64.whl#sha256=71ad2f82da0f41eaec159f39fc85854e27c2391efa91b373e550648a6f4aaad3 ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp312-cp312-linux_x86_64.whl#sha256=b473571d478912f92881cc13f15fa18f8463fb0fb8a068c96ed47a7d45a4da0a ; platform_system == 'Linux' and python_version == '3.12' and platform_machine == 'x86_64'",
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp313-cp313-linux_x86_64.whl#sha256=3bc64a746ff25a93de140902c60c9e819d7413f5cea1e88d80999c27a5901e9c ; platform_system == 'Linux' and python_version == '3.13' and platform_machine == 'x86_64'",
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp310-cp310-win_amd64.whl#sha256=ce50691ab3fb6301d9b7bb8b3834cf5fa7152a2b5f91fd24c5efdc601a25b780 ; sys_platform == 'win32' and python_version == '3.10' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp311-cp311-win_amd64.whl#sha256=cb9d37f21cb9fb7df67d62863f021c3144e8d8832b9ea8e8523ac308bc620ea1 ; sys_platform == 'win32' and python_version == '3.11' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp312-cp312-win_amd64.whl#sha256=3ad605be4728b6d3a28a44d07dd794b1a9e45551b0057815bf25eb2a6d6a56a7 ; sys_platform == 'win32' and python_version == '3.12' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "torch @ https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp313-cp313-win_amd64.whl#sha256=2b4b56dd6c792aef82006904fa888692e3782e4ae5da27526801bad4898f05a5 ; sys_platform == 'win32' and python_version == '3.13' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+
+ "bitsandbytes @ https://github.com/bitsandbytes-foundation/bitsandbytes/releases/download/continuous-release_main/bitsandbytes-1.33.7.preview-py3-none-manylinux_2_24_x86_64.whl ; ('linux' in sys_platform) and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "bitsandbytes @ https://github.com/bitsandbytes-foundation/bitsandbytes/releases/download/continuous-release_main/bitsandbytes-1.33.7.preview-py3-none-win_amd64.whl ; (sys_platform == 'win32') and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp310-cp310-manylinux_2_28_x86_64.whl#sha256=7e1e7b170fcf7161c8499b67156c5a05462243626dc0974010791a0bab4378d3 ; platform_system == 'Linux' and python_version == '3.10' and platform_machine == 'x86_64'",
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp311-cp311-manylinux_2_28_x86_64.whl#sha256=bd6add201bd7628af70437292e1447abb368e0b5f4ff9abd334ae435efd44792 ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp312-cp312-manylinux_2_28_x86_64.whl#sha256=6ad2543496bc29e59d3dd614a94d09aa9870318aedb66045344fffddfedd2cf8 ; platform_system == 'Linux' and python_version == '3.12' and platform_machine == 'x86_64'",
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp313-cp313-manylinux_2_28_x86_64.whl#sha256=80269f37865fcd8b57f20e4786efae2200bfa2b2727926c3c7acc82f0e7d3548 ; platform_system == 'Linux' and python_version == '3.13' and platform_machine == 'x86_64'",
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp310-cp310-win_amd64.whl#sha256=6b9485ba85dcba4d196d6134d9c3332fb228fb2556416bf0450a64e8a472fcba ; sys_platform == 'win32' and python_version == '3.10' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp311-cp311-win_amd64.whl#sha256=36cbaedf10f6412af5c89afd9aeea474e6a56a0050348ada8fabe1ecaf6b879e ; sys_platform == 'win32' and python_version == '3.11' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp312-cp312-win_amd64.whl#sha256=738357d97468d75fe3d510ac37e65130f2787f81d9bbc1518898f7396dc3403f ; sys_platform == 'win32' and python_version == '3.12' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp313-cp313-win_amd64.whl#sha256=1c4b44b36a557f7381e3076fb8843366742238648441d607c8d049c6da0f8886 ; sys_platform == 'win32' and python_version == '3.13' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+]
+intel-gpu-torch210 = [
+ "unsloth[intelgputorch210]"
+]
intel = [
"unsloth[intelgputorch280]",
]
amd = [
"unsloth[huggingfacenotorch]",
- "bitsandbytes @ https://github.com/bitsandbytes-foundation/bitsandbytes/releases/download/continuous-release_main/bitsandbytes-1.33.7.preview-py3-none-manylinux_2_24_x86_64.whl ; ('linux' in sys_platform) and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
- "bitsandbytes @ https://github.com/bitsandbytes-foundation/bitsandbytes/releases/download/continuous-release_main/bitsandbytes-1.33.7.preview-py3-none-win_amd64.whl ; (sys_platform == 'win32') and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
- "bitsandbytes @ https://github.com/bitsandbytes-foundation/bitsandbytes/releases/download/continuous-release_main/bitsandbytes-1.33.7.preview-py3-none-manylinux_2_24_aarch64.whl ; ('linux' in sys_platform) and (platform_machine == 'aarch64')",
+ "bitsandbytes>=0.49.1 ; ('linux' in sys_platform) and (platform_machine == 'AMD64' or platform_machine == 'x86_64' or platform_machine == 'aarch64')",
+ "bitsandbytes>=0.49.1 ; (sys_platform == 'win32') and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
]
[project.urls]
diff --git a/tests/test_get_model_name.py b/tests/test_get_model_name.py
new file mode 100644
index 0000000000..ad89f595f0
--- /dev/null
+++ b/tests/test_get_model_name.py
@@ -0,0 +1,127 @@
+import unittest
+from unittest.mock import patch
+from unsloth.models.loader_utils import get_model_name
+from unsloth.models import loader_utils
+from unsloth.models.mapper import FLOAT_TO_INT_MAPPER, MAP_TO_UNSLOTH_16bit
+
+
+def _no_remote_mapper():
+ return {}, {}, {}
+
+
+class TestGetModelName(unittest.TestCase):
+ def _assert_mapping(self, model_name, load_in_4bit, expected, should_change):
+ mapped = get_model_name(model_name, load_in_4bit = load_in_4bit)
+ self.assertEqual(mapped.lower(), expected.lower())
+ if should_change:
+ self.assertNotEqual(mapped.lower(), model_name.lower())
+ else:
+ self.assertEqual(mapped.lower(), model_name.lower())
+
+ @patch.object(loader_utils, "_get_new_mapper", _no_remote_mapper)
+ def test_resolution_matrix(self):
+ cases = [
+ # Core mappings
+ ("meta-llama/Llama-2-7b-hf", True, "unsloth/llama-2-7b-bnb-4bit", True),
+ ("meta-llama/Llama-2-7b-hf", False, "unsloth/llama-2-7b", True),
+ (
+ "mistralai/Ministral-8B-Instruct-2410",
+ True,
+ "mistralai/Ministral-8B-Instruct-2410",
+ False,
+ ),
+ (
+ "meta-llama/Llama-3.2-1B-Instruct",
+ False,
+ "unsloth/Llama-3.2-1B-Instruct",
+ True,
+ ),
+ (
+ "meta-llama/Llama-2-7b-chat-hf",
+ True,
+ "unsloth/llama-2-7b-chat-bnb-4bit",
+ True,
+ ),
+ (
+ "meta-llama/Llama-3.3-70B-Instruct",
+ True,
+ "unsloth/llama-3.3-70b-instruct-unsloth-bnb-4bit",
+ True,
+ ),
+ ("Qwen/Qwen3-8B", True, "unsloth/Qwen3-8B-unsloth-bnb-4bit", True),
+ ("Qwen/Qwen3-8B", False, "unsloth/Qwen3-8B", True),
+ ("Qwen/Qwen3-8B-FP8", False, "unsloth/Qwen3-8B-FP8", True),
+ ("Qwen/Qwen3-8B-FP8", True, "unsloth/Qwen3-8B-unsloth-bnb-4bit", True),
+ (
+ "mistralai/Ministral-3-3B-Instruct-2512",
+ True,
+ "unsloth/Ministral-3-3B-Instruct-2512-unsloth-bnb-4bit",
+ True,
+ ),
+ (
+ "mistralai/Ministral-3-3B-Instruct-2512",
+ False,
+ "unsloth/Ministral-3-3B-Instruct-2512",
+ True,
+ ),
+ ("unsloth/Kimi-K2-Instruct", True, "unsloth/Kimi-K2-Instruct-BF16", True),
+ ("unsloth/Kimi-K2-Instruct", False, "unsloth/Kimi-K2-Instruct", False),
+ # Fallback-to-original behavior
+ "nonexistent-user/nonexistent-model-123",
+ "google/gemma-3-random-prototype-123",
+ "imdatta0/nanoqwen-fp8",
+ "imdatta0/nanoqwen-bf16",
+ # Backward compatibility for legacy 4bit names
+ ("unsloth/llama-2-7b-bnb-4bit", True, "unsloth/llama-2-7b-bnb-4bit", False),
+ ("unsloth/llama-2-7b-bnb-4bit", False, "unsloth/llama-2-7b", True),
+ ("google/gemma-2-9b", True, "unsloth/gemma-2-9b-bnb-4bit", True),
+ # GPT-OSS behavior
+ ("openai/gpt-oss-20b", False, "unsloth/gpt-oss-20b", True),
+ ("openai/gpt-oss-20b", True, "unsloth/gpt-oss-20b-unsloth-bnb-4bit", True),
+ ("unsloth/gpt-oss-20b", True, "unsloth/gpt-oss-20b-unsloth-bnb-4bit", True),
+ ("unsloth/gpt-oss-20b-bf16", True, "unsloth/gpt-oss-20b-bf16", False),
+ (
+ "unsloth/gpt-oss-20b-unsloth-bnb-4bit",
+ False,
+ "unsloth/gpt-oss-20b",
+ True,
+ ),
+ (
+ "unsloth/gpt-oss-20b-bnb-4bit",
+ True,
+ "unsloth/gpt-oss-20b-bnb-4bit",
+ False,
+ ),
+ ]
+ for case in cases:
+ if isinstance(case, str):
+ model_name = case
+ with self.subTest(model_name = model_name, load_in_4bit = True):
+ self._assert_mapping(model_name, True, model_name, False)
+ else:
+ model_name, load_in_4bit, expected, should_change = case
+ with self.subTest(model_name = model_name, load_in_4bit = load_in_4bit):
+ self._assert_mapping(
+ model_name, load_in_4bit, expected, should_change
+ )
+
+ def test_static_mapper_contract(self):
+ contracts = [
+ ("qwen/qwen3-8b", "unsloth/qwen3-8b-unsloth-bnb-4bit"),
+ ("qwen/qwen3-8b-fp8", "unsloth/qwen3-8b-unsloth-bnb-4bit"),
+ (
+ "mistralai/ministral-3-3b-instruct-2512",
+ "unsloth/ministral-3-3b-instruct-2512-unsloth-bnb-4bit",
+ ),
+ ("unsloth/kimi-k2-instruct", "unsloth/kimi-k2-instruct-bf16"),
+ ]
+ for src, expected in contracts:
+ with self.subTest(src = src):
+ self.assertEqual(FLOAT_TO_INT_MAPPER[src], expected)
+ self.assertEqual(
+ MAP_TO_UNSLOTH_16bit["qwen/qwen3-8b-fp8"], "unsloth/Qwen3-8B-FP8"
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/tests/utils/test_trunc_normal_patch.py b/tests/utils/test_trunc_normal_patch.py
new file mode 100644
index 0000000000..b84a0772d8
--- /dev/null
+++ b/tests/utils/test_trunc_normal_patch.py
@@ -0,0 +1,114 @@
+# 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 Lesser 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 General Public License for more details.
+#
+# You should have received a copy of the GNU Lesser General Public License
+# along with this program. If not, see .
+
+"""Tests for trunc_normal low-precision patch compatibility."""
+
+import importlib.util
+import inspect
+from pathlib import Path
+
+import pytest
+import torch
+
+
+_MISSING = object()
+
+
+def _load_import_fixes_module():
+ repo_root = Path(__file__).resolve().parents[2]
+ import_fixes_path = repo_root / "unsloth" / "import_fixes.py"
+ spec = importlib.util.spec_from_file_location(
+ "unsloth_import_fixes_local", import_fixes_path
+ )
+ assert spec is not None and spec.loader is not None
+ module = importlib.util.module_from_spec(spec)
+ spec.loader.exec_module(module)
+ return module
+
+
+def _getattr_or_missing(obj, name):
+ return getattr(obj, name) if hasattr(obj, name) else _MISSING
+
+
+def _restore_attr(obj, name, value):
+ if value is _MISSING:
+ if hasattr(obj, name):
+ delattr(obj, name)
+ return
+ setattr(obj, name, value)
+
+
+def test_trunc_normal_patch_accepts_positional_generator():
+ import_fixes = _load_import_fixes_module()
+ patch_fn = import_fixes.patch_trunc_normal_precision_issue
+
+ init_mod = torch.nn.init
+ old_fn = init_mod.trunc_normal_
+ old_patched = _getattr_or_missing(init_mod, "_unsloth_trunc_normal_patched")
+ old_original = _getattr_or_missing(init_mod, "_unsloth_trunc_normal_original")
+ try:
+ # Normalize to an unpatched baseline before applying the patch.
+ if old_original is not _MISSING:
+ init_mod.trunc_normal_ = old_original
+ if hasattr(init_mod, "_unsloth_trunc_normal_patched"):
+ delattr(init_mod, "_unsloth_trunc_normal_patched")
+ if hasattr(init_mod, "_unsloth_trunc_normal_original"):
+ delattr(init_mod, "_unsloth_trunc_normal_original")
+
+ patch_fn()
+ sig = inspect.signature(init_mod.trunc_normal_)
+ assert "generator" in sig.parameters
+ assert sig.parameters["generator"].kind is not inspect.Parameter.KEYWORD_ONLY
+
+ tensor = torch.empty(1024, dtype = torch.float32)
+ gen = torch.Generator()
+ gen.manual_seed(3407)
+
+ init_mod.trunc_normal_(tensor, 0.0, 1.0, -2.0, 2.0, gen)
+ init_mod.trunc_normal_(tensor, mean = 0.0, std = 1.0, a = -2.0, b = 2.0, generator = gen)
+ finally:
+ init_mod.trunc_normal_ = old_fn
+ _restore_attr(init_mod, "_unsloth_trunc_normal_patched", old_patched)
+ _restore_attr(init_mod, "_unsloth_trunc_normal_original", old_original)
+
+
+def test_trunc_normal_patch_rejects_invalid_generator():
+ import_fixes = _load_import_fixes_module()
+ patch_fn = import_fixes.patch_trunc_normal_precision_issue
+
+ init_mod = torch.nn.init
+ old_fn = init_mod.trunc_normal_
+ old_patched = _getattr_or_missing(init_mod, "_unsloth_trunc_normal_patched")
+ old_original = _getattr_or_missing(init_mod, "_unsloth_trunc_normal_original")
+ try:
+ if old_original is not _MISSING:
+ init_mod.trunc_normal_ = old_original
+ if hasattr(init_mod, "_unsloth_trunc_normal_patched"):
+ delattr(init_mod, "_unsloth_trunc_normal_patched")
+ if hasattr(init_mod, "_unsloth_trunc_normal_original"):
+ delattr(init_mod, "_unsloth_trunc_normal_original")
+
+ patch_fn()
+ sig = inspect.signature(init_mod.trunc_normal_)
+ if "generator" not in sig.parameters:
+ pytest.skip("torch.nn.init.trunc_normal_ lacks a generator parameter")
+
+ tensor = torch.empty(16, dtype = torch.float32)
+ with pytest.raises(TypeError):
+ init_mod.trunc_normal_(tensor, generator = 123)
+ finally:
+ init_mod.trunc_normal_ = old_fn
+ _restore_attr(init_mod, "_unsloth_trunc_normal_patched", old_patched)
+ _restore_attr(init_mod, "_unsloth_trunc_normal_original", old_original)
diff --git a/unsloth/__init__.py b/unsloth/__init__.py
index b068d6a5fc..aff9e3d186 100644
--- a/unsloth/__init__.py
+++ b/unsloth/__init__.py
@@ -28,16 +28,26 @@ already_imported = [mod for mod in critical_modules if mod in sys.modules]
from .import_fixes import (
fix_message_factory_issue,
check_fbgemm_gpu_version,
+ disable_broken_causal_conv1d,
+ disable_broken_vllm,
+ configure_amdgpu_asic_id_table_path,
torchvision_compatibility_check,
fix_diffusers_warnings,
fix_huggingface_hub,
)
+# Configure libdrm ids table path early so ROCm can resolve AMD GPU names.
+configure_amdgpu_asic_id_table_path()
+disable_broken_causal_conv1d()
+disable_broken_vllm()
fix_message_factory_issue()
check_fbgemm_gpu_version()
torchvision_compatibility_check()
fix_diffusers_warnings()
fix_huggingface_hub()
+del configure_amdgpu_asic_id_table_path
+del disable_broken_causal_conv1d
+del disable_broken_vllm
del fix_message_factory_issue
del check_fbgemm_gpu_version
del torchvision_compatibility_check
@@ -79,7 +89,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.2.1"):
+ if Version(unsloth_zoo_version) < Version("2026.3.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`"
@@ -129,7 +139,7 @@ from .import_fixes import (
fix_vllm_guided_decoding_params,
fix_vllm_pdl_blackwell,
fix_triton_compiled_kernel_missing_attrs,
- fix_rocm_triton_key_error,
+ patch_trunc_normal_precision_issue,
ignore_logger_messages,
patch_ipykernel_hf_xet,
patch_trackio,
@@ -150,7 +160,7 @@ check_vllm_torch_sm100_compatibility()
fix_vllm_guided_decoding_params()
fix_vllm_pdl_blackwell()
fix_triton_compiled_kernel_missing_attrs()
-fix_rocm_triton_key_error()
+patch_trunc_normal_precision_issue()
ignore_logger_messages()
patch_ipykernel_hf_xet()
patch_trackio()
@@ -169,7 +179,7 @@ del check_vllm_torch_sm100_compatibility
del fix_vllm_guided_decoding_params
del fix_vllm_pdl_blackwell
del fix_triton_compiled_kernel_missing_attrs
-del fix_rocm_triton_key_error
+del patch_trunc_normal_precision_issue
del ignore_logger_messages
del patch_ipykernel_hf_xet
del patch_trackio
diff --git a/unsloth/chat_templates.py b/unsloth/chat_templates.py
index a17a6f6299..35eb871529 100644
--- a/unsloth/chat_templates.py
+++ b/unsloth/chat_templates.py
@@ -1959,41 +1959,61 @@ def _parse_combined_prompt(combined_prompt, dataset):
def _create_formatter(possible_columns, final_optional_prompts, user_column_name):
- # Start final prompt!
- function = ["def __combined_prompt_processor__(examples):"]
- columns = list(set(possible_columns))
- for column in columns:
- function.append(f"{' '*4}{column}__ = examples['{column}']")
- function.append(f"{' '*4}texts = []")
- function.append(f"{' '*4}for ({', '.join(columns)}) in zip({', '.join(f'{x}__' for x in columns)}):")
-
- # Add optional tags as well!
- final_prompt = ""
- formatter = []
+ columns = list(dict.fromkeys(possible_columns))
+ merged_prompt_parts = []
+ formatter_templates = []
for j, optional_prompt in enumerate(final_optional_prompts):
if type(optional_prompt) is str:
- columns = re.findall(r"\{(.+?)\}", optional_prompt)
- formatter += columns
- # Must escape \n \r
- final_prompt += optional_prompt.encode("unicode-escape").decode("utf-8").replace("'", "\\'").replace('"', '\\"')
- else:
- where, prompt = optional_prompt
- # Strip [[...]]
- # Must escape \n \r
- prompt = prompt[2:-2].encode("unicode-escape").decode("utf-8").replace("'", "\\'").replace('"', '\\"')
- columns = re.findall(r"\{(.+?)\}", prompt)
- x = f"__optional_{j}__"
- prompt = f"{' '*8}{x} = '{prompt}'.format({', '.join(f'{x} = {x}' for x in columns)}) if {columns[0]} else ''"
- function.append(prompt)
- formatter.append(x)
- final_prompt += "{" + x + "}"
+ needed_columns = re.findall(r"\{(.+?)\}", optional_prompt)
+ formatter_templates.append(("required", optional_prompt, needed_columns))
+ merged_prompt_parts.append(optional_prompt)
+ continue
- function.insert(1, f"{' '*4}__combined_prompt__ = '{final_prompt}'")
- function.append(f"{' '*8}texts.append("\
- f"__combined_prompt__.format({', '.join(f'{x} = {x}' for x in formatter)}))")
- function.append(f"{' '*4}return " + "{ " + f"'{user_column_name}' : texts" + " }")
- return "\n".join(function)
+ _, prompt = optional_prompt
+ prompt = prompt[2:-2]
+ needed_columns = re.findall(r"\{(.+?)\}", prompt)
+ if len(needed_columns) == 0:
+ raise IndexError("Unsloth: Optional [[...]] blocks must contain at least 1 {column}.")
+ optional_name = f"__optional_{j}__"
+ formatter_templates.append(("optional", optional_name, prompt, needed_columns))
+ merged_prompt_parts.append("{" + optional_name + "}")
+
+ merged_prompt = "".join(merged_prompt_parts)
+
+ def __combined_prompt_processor__(examples):
+ if len(examples) == 0:
+ return {user_column_name: []}
+
+ first_key = next(iter(examples.keys()), None)
+ if first_key is None:
+ return {user_column_name: []}
+ n_rows = len(examples[first_key])
+
+ texts = []
+ for row_idx in range(n_rows):
+ row_values = {column: examples[column][row_idx] for column in columns}
+ formatter_values = {}
+
+ for formatter_template in formatter_templates:
+ if formatter_template[0] == "required":
+ _, _, needed_columns = formatter_template
+ for column in needed_columns:
+ formatter_values[column] = row_values[column]
+ continue
+
+ _, optional_name, prompt, needed_columns = formatter_template
+ if row_values[needed_columns[0]] not in (None, ""):
+ prompt_values = {column: row_values[column] for column in needed_columns}
+ formatter_values[optional_name] = prompt.format(**prompt_values)
+ else:
+ formatter_values[optional_name] = ""
+
+ texts.append(merged_prompt.format(**formatter_values))
+
+ return {user_column_name: texts}
+
+ return __combined_prompt_processor__
def to_sharegpt(
@@ -2025,13 +2045,17 @@ def to_sharegpt(
raise TypeError("Unsloth: Your dataset is probably already in ShareGPT format!")
possible_columns, final_optional_prompts = _parse_combined_prompt(merged_prompt, dataset)
- function = _create_formatter(possible_columns, final_optional_prompts, merged_column_name)
- exec(function, globals())
- dataset = dataset.map(__combined_prompt_processor__, batched = True, desc = "Merging columns")
+ formatter = _create_formatter(possible_columns, final_optional_prompts, merged_column_name)
+ dataset = dataset.map(formatter, batched = True, desc = "Merging columns")
def __convert_to_sharegpt__(examples):
users = examples[merged_column_name]
assistants = examples[output_column_name]
+ if len(users) != len(assistants):
+ raise ValueError(
+ "Unsloth: Input and output columns must have matching batch lengths. "
+ f"Got {len(users)} {merged_column_name} rows and {len(assistants)} {output_column_name} rows."
+ )
texts = [
[
{"from" : "human", "value" : str(user) },
@@ -2062,19 +2086,18 @@ def to_sharegpt(
dataset = concatenate_datasets(all_shuffled, axis = 1)
# Combine them into 1
- function = "def __combine_conversations__(examples):\n"
n_extensions += 1
- for j in range(n_extensions):
- function += f"{' '*4}conversations{j}__ = examples['conversations{j}']\n"
- function += f"{' '*4}convos = []\n"
- function += f"{' '*4}for ({', '.join(f'conversations{j}' for j in range(n_extensions))}) "\
- f"in zip({', '.join(f'conversations{j}__' for j in range(n_extensions))}):\n"
- function += f"{' '*8}convos.append("\
- f"{'+'.join(f'conversations{j}' for j in range(n_extensions))})\n"
- function += f"{' '*4}return " + "{ " + "'conversations' : convos" + " }"
+ conversation_columns = [f"conversations{j}" for j in range(n_extensions)]
+ def __combine_conversations__(examples):
+ columns = [examples[column] for column in conversation_columns]
+ convos = []
+ for conversations in zip(*columns):
+ merged_conversation = []
+ for conversation in conversations:
+ merged_conversation.extend(conversation)
+ convos.append(merged_conversation)
+ return {"conversations" : convos}
- # Map function
- exec(function, globals())
dataset = dataset.map(
__combine_conversations__,
batched = True,
@@ -2682,16 +2705,23 @@ def test_hf_gguf_equivalence(tokenizer, gguf_model = "./model-unsloth.F16.gguf")
if tokenizer.chat_template is not None:
prompt = tokenizer.apply_chat_template(messages, tokenize = False, add_generation_prompt = True)
- prompt = prompt.replace("'", "") # Subprocess does not like ''
prompt = remove_special_tokens(tokenizer, prompt)
prompts.append(prompt)
for prompt in prompts:
- command = f"./llama.cpp/llama-cli -m {gguf_model} -n 0 --temp 0.0 --verbose-prompt "\
- f"--check-tensors -p '{prompt}'"
+ # Use a list of args with shell=False so prompt content is passed literally.
+ command = [
+ "./llama.cpp/llama-cli",
+ "-m", gguf_model,
+ "-n", "0",
+ "--temp", "0.0",
+ "--verbose-prompt",
+ "--check-tensors",
+ "-p", prompt,
+ ]
datas = []
- with subprocess.Popen(command, shell = True, stdout = subprocess.PIPE, stderr = subprocess.STDOUT, bufsize = 1) as sp:
+ with subprocess.Popen(command, shell = False, stdout = subprocess.PIPE, stderr = subprocess.STDOUT, bufsize = 1) as sp:
for line in sp.stdout:
datas.append(line.decode("utf-8", errors = "replace"))
gguf_tokens = "".join(datas)
diff --git a/unsloth/dataprep/synthetic.py b/unsloth/dataprep/synthetic.py
index 7735b1ecd2..612c531f47 100644
--- a/unsloth/dataprep/synthetic.py
+++ b/unsloth/dataprep/synthetic.py
@@ -27,11 +27,6 @@ import torch
import gc
import time
import re
-from unsloth_zoo.vllm_utils import (
- load_vllm,
- patch_vllm,
- delete_vllm,
-)
from unsloth_zoo.log import logger
import numpy as np
@@ -40,6 +35,16 @@ from .synthetic_configs import (
)
+def _load_vllm_utils():
+ from unsloth_zoo.vllm_utils import (
+ load_vllm,
+ patch_vllm,
+ delete_vllm,
+ )
+
+ return load_vllm, patch_vllm, delete_vllm
+
+
def terminate_tree(proc: subprocess.Popen, timeout = 15):
if proc is None or proc.poll() is not None:
return
@@ -182,6 +187,8 @@ class SyntheticDataKit:
model_name,
token = token,
)
+ load_vllm, patch_vllm, delete_vllm = _load_vllm_utils()
+ self._delete_vllm = delete_vllm
patch_vllm(debug = False)
engine_args = load_vllm(
model_name = model_name,
@@ -364,7 +371,8 @@ class SyntheticDataKit:
gc.collect()
# Delete vLLM module as well
- delete_vllm(llm = None)
+ if hasattr(self, "_delete_vllm"):
+ self._delete_vllm(llm = None)
def __enter__(self):
return self
diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py
index 4a2211f9a0..ebd81f9568 100644
--- a/unsloth/import_fixes.py
+++ b/unsloth/import_fixes.py
@@ -13,6 +13,8 @@
# limitations under the License.
import os
+import importlib.abc
+import importlib.machinery
import importlib.util
from pathlib import Path
from importlib.metadata import version as importlib_version
@@ -21,6 +23,8 @@ import re
import logging
import textwrap
import warnings
+import sys
+import functools
# We cannot do from unsloth_zoo.log import logger since FBGEMM might cause seg faults.
UNSLOTH_ENABLE_LOGGING = os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") in (
@@ -40,6 +44,8 @@ else:
)
logger.setLevel(logging.WARNING)
+_AMDGPU_IDS_MISSING_TEXT = "amdgpu.ids: No such file or directory"
+
def Version(version):
try:
@@ -91,7 +97,54 @@ class HidePrintMessage:
return getattr(self._original_stream, name)
-if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1":
+import contextlib
+import ctypes
+
+try:
+ _libc = ctypes.CDLL(None)
+except Exception:
+ _libc = None
+
+
+@contextlib.contextmanager
+def suppress_cuda_printf():
+ """Suppress CUDA device-side printf by redirecting stdout/stderr fds to /dev/null.
+
+ CUDA device printf (eg CUTLASS "Arch conditional MMA" errors on Blackwell)
+ writes to stdout fd 1 at the C level, bypassing Python sys.stdout entirely.
+ The existing HidePrintMessage filter on sys.stderr cannot catch these since
+ they go to a different fd at a different layer. This context manager redirects
+ both fd 1 and fd 2 at the OS level, syncs CUDA, then restores them.
+ """
+ sys.stdout.flush()
+ sys.stderr.flush()
+ saved_fds = {}
+ try:
+ for fd in (1, 2):
+ saved_fds[fd] = os.dup(fd)
+ devnull = os.open(os.devnull, os.O_WRONLY)
+ os.dup2(devnull, fd)
+ os.close(devnull)
+ yield
+ finally:
+ try:
+ import torch
+
+ if torch.cuda.is_available():
+ torch.cuda.synchronize()
+ except Exception:
+ pass
+ if _libc is not None:
+ try:
+ _libc.fflush(None)
+ except Exception:
+ pass
+ for fd, saved in saved_fds.items():
+ os.dup2(saved, fd)
+ os.close(saved)
+
+
+if not UNSLOTH_ENABLE_LOGGING:
import sys
# Apply to stderr for FBGEMM and CUTLASS errors
@@ -196,6 +249,15 @@ if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1":
# Triton "df: No such file or directory" stderr noise
sys.stderr.add_filter("df: No such file")
+ # ROCm/libdrm missing ids table stderr noise on some AMD setups
+ sys.stderr.add_filter(_AMDGPU_IDS_MISSING_TEXT)
+ # Apex ROCm fused RoPE backend selection warning when Aiter is enabled.
+ warnings.filterwarnings(
+ "ignore",
+ message = r"^Aiter backend is selected for fused RoPE\.?",
+ category = UserWarning,
+ module = r"^apex\.transformer\.functional\.fused_rope$",
+ )
# Fix up AttributeError: 'MessageFactory' object has no attribute 'GetPrototype'
@@ -406,14 +468,18 @@ def fix_vllm_guided_decoding_params():
# trl still wants to use GuidedDecodingParams. This is a temporary patch till trl updates
try:
import vllm
- except ImportError as e:
+ except (ImportError, OSError) as e:
_maybe_raise_vllm_transformers_mismatch(e)
+ if disable_broken_vllm(e):
+ return
raise
try:
from vllm.sampling_params import GuidedDecodingParams
- except ImportError as e:
+ except (ImportError, OSError) as e:
_maybe_raise_vllm_transformers_mismatch(e)
+ if disable_broken_vllm(e):
+ return
if not hasattr(vllm, "sampling_params") or not hasattr(
vllm.sampling_params, "StructuredOutputsParams"
):
@@ -653,7 +719,6 @@ def torchvision_compatibility_check():
# 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)
@@ -675,6 +740,7 @@ def torchvision_compatibility_check():
f"Unsloth: torch=={torch_version_raw} requires "
f"torchvision>={required_tv_str}, "
f"but found torchvision=={torchvision_version_raw}. "
+ f'Try updating torchvision via `pip install --upgrade "torchvision>={required_tv_str}"`. '
f"Please refer to https://pytorch.org/get-started/previous-versions/ "
f"for more information."
)
@@ -690,15 +756,10 @@ def torchvision_compatibility_check():
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"
- )
+ # Only downgrade to warning for custom/source or prerelease builds.
+ # Stable mismatches should fail fast to prevent runtime operator errors.
+ if is_custom or is_prerelease:
+ reason = "custom/source build" if is_custom else "pre-release build"
logger.warning(
f"{message}\n"
f"Detected a {reason}. "
@@ -880,37 +941,86 @@ def fix_triton_compiled_kernel_missing_attrs():
)
-def fix_rocm_triton_key_error():
+def patch_trunc_normal_precision_issue():
"""
- ROCm + torch.compile can fail if Triton lacks `triton_key`.
- Disable Inductor/compile only on ROCm when that symbol is missing.
+ Patch torch.nn.init.trunc_normal_ for low precision tensors to run init in fp32.
+
+ torch.nn.init.trunc_normal_ can saturate at truncation bounds in fp16/bf16 on
+ some versions/backends. This was observed in TorchTitan investigations where
+ low-precision truncation produced boundary-heavy initialization behavior:
+ https://github.com/pytorch/torchtitan/pull/2342
+
+ To avoid that failure mode, initialize into a temporary fp32 tensor, then copy
+ back to the original dtype.
"""
try:
import torch
except (ImportError, ModuleNotFoundError):
return
- if not getattr(torch.version, "hip", None):
+ if getattr(torch.nn.init, "_unsloth_trunc_normal_patched", False):
return
+ original_trunc_normal = torch.nn.init.trunc_normal_
+ if getattr(original_trunc_normal, "__unsloth_trunc_normal_patched__", False):
+ torch.nn.init._unsloth_trunc_normal_patched = True
+ return
+
+ low_precision_dtypes = {torch.float16, torch.bfloat16}
+
+ def _call_original(target, mean, std, a, b, generator):
+ if generator is None:
+ return original_trunc_normal(target, mean = mean, std = std, a = a, b = b)
+ try:
+ return original_trunc_normal(
+ target, mean = mean, std = std, a = a, b = b, generator = generator
+ )
+ except TypeError as exc:
+ # Older torch versions may not accept a generator keyword argument.
+ msg = str(exc).lower()
+ if "unexpected keyword argument" in msg and "generator" in msg:
+ return original_trunc_normal(target, mean = mean, std = std, a = a, b = b)
+ raise
+
try:
- import triton
- except (ImportError, ModuleNotFoundError):
- return
+ from torch.distributed._tensor import DTensor
+ except Exception:
+ DTensor = None
- try:
- from triton.runtime import triton_key # noqa: F401
+ @torch.no_grad()
+ def _patched_trunc_normal_(
+ tensor,
+ mean: float = 0.0,
+ std: float = 1.0,
+ a: float = -2.0,
+ b: float = 2.0,
+ generator = None,
+ ):
+ if DTensor is not None and isinstance(tensor, DTensor):
+ local_tensor = getattr(tensor, "_local_tensor", None)
+ if local_tensor is None:
+ return _call_original(tensor, mean, std, a, b, generator)
+ if local_tensor.dtype in low_precision_dtypes:
+ local_fp32 = local_tensor.float()
+ _call_original(local_fp32, mean, std, a, b, generator)
+ local_tensor.copy_(local_fp32.to(dtype = local_tensor.dtype))
+ return tensor
+ return _call_original(tensor, mean, std, a, b, generator)
- return
- except ImportError:
- pass
+ if tensor.dtype in low_precision_dtypes:
+ tensor_fp32 = tensor.float()
+ _call_original(tensor_fp32, mean, std, a, b, generator)
+ tensor.copy_(tensor_fp32.to(dtype = tensor.dtype))
+ return tensor
- 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."
- )
+ return _call_original(tensor, mean, std, a, b, generator)
+
+ _patched_trunc_normal_.__unsloth_trunc_normal_patched__ = True
+ _patched_trunc_normal_._unsloth_original = original_trunc_normal
+ torch.nn.init._unsloth_trunc_normal_original = original_trunc_normal
+ torch.nn.init.trunc_normal_ = _patched_trunc_normal_
+ torch.nn.init._unsloth_trunc_normal_patched = True
+ logger.info("Unsloth: Patched torch.nn.init.trunc_normal_ for fp16/bf16 stability.")
def check_vllm_torch_sm100_compatibility():
@@ -1018,7 +1128,7 @@ def fix_vllm_pdl_blackwell():
def _spec_exists(name):
try:
return importlib.util.find_spec(name) is not None
- except (ModuleNotFoundError, ValueError):
+ except (ImportError, OSError, ModuleNotFoundError, ValueError):
return False
# Check if vLLM has the PDL-related modules before doing internet check
@@ -1031,10 +1141,10 @@ def fix_vllm_pdl_blackwell():
return
# Check if vLLM version includes the fix
- VLLM_PDL_FIX_VERSION = "0.13.2"
+ VLLM_PDL_FIX_VERSION = "0.15.0"
try:
vllm_version = Version(importlib_version("vllm"))
- if vllm_version > Version(VLLM_PDL_FIX_VERSION):
+ if vllm_version >= Version(VLLM_PDL_FIX_VERSION):
logger.info(
f"Unsloth: SM100 ({sm100_gpu_name}) detected but vLLM {vllm_version} "
f"should include PDL fix - skipping workaround"
@@ -1052,6 +1162,12 @@ def fix_vllm_pdl_blackwell():
return False
patched = []
+ patched_names = set()
+
+ def _record_patch(name):
+ if name not in patched_names:
+ patched.append(name)
+ patched_names.add(name)
# First, patch the source module (utils.py) where supports_pdl is defined.
# This is critical because supports_pdl uses @lru_cache - we must clear the
@@ -1063,7 +1179,7 @@ def fix_vllm_pdl_blackwell():
if hasattr(original_fn, "cache_clear"):
original_fn.cache_clear()
utils_module.supports_pdl = fake_supports_pdl
- patched.append("utils")
+ _record_patch("utils")
except (ImportError, ModuleNotFoundError, AttributeError):
pass
@@ -1080,10 +1196,19 @@ def fix_vllm_pdl_blackwell():
module = importlib.import_module(path)
if hasattr(module, "supports_pdl"):
module.supports_pdl = fake_supports_pdl
- patched.append(name)
+ _record_patch(name)
except (ImportError, ModuleNotFoundError, AttributeError):
pass
+ # Patch any additional already-loaded triton ops consumers that expose supports_pdl.
+ for module_name, module in tuple(sys.modules.items()):
+ if not module_name.startswith("vllm.lora.ops.triton_ops."):
+ continue
+ if module is None or not hasattr(module, "supports_pdl"):
+ continue
+ module.supports_pdl = fake_supports_pdl
+ _record_patch(module_name.rsplit(".", 1)[-1])
+
if patched:
logger.info(
f"Unsloth: Applied PDL fix for SM100 ({sm100_gpu_name}) - "
@@ -1163,3 +1288,481 @@ def disable_torchcodec_if_broken():
tf_import_utils._torchcodec_available = False
except (ImportError, AttributeError):
pass
+
+
+CAUSAL_CONV1D_BROKEN = False
+_CAUSAL_CONV1D_PREFIX = "causal_conv1d"
+_CAUSAL_CONV1D_BLOCKER_SENTINEL = "_unsloth_causal_conv1d_blocker"
+VLLM_BROKEN = False
+_VLLM_PREFIX = "vllm"
+_VLLM_BLOCKER_SENTINEL = "_unsloth_vllm_blocker"
+_ROCM_ENV_HINT_KEYS = (
+ "ROCM_PATH",
+ "ROCM_HOME",
+ "HIP_PATH",
+ "HSA_PATH",
+ "HIP_VISIBLE_DEVICES",
+ "ROCR_VISIBLE_DEVICES",
+)
+_ROCM_PATH_HINTS = (
+ Path("/opt/rocm"),
+ Path("/dev/kfd"),
+ Path("/sys/module/amdgpu"),
+)
+_AMDGPU_ASIC_ID_TABLE_PATH_ENV = "AMDGPU_ASIC_ID_TABLE_PATH"
+_AMDGPU_ASIC_ID_CANDIDATE_PATHS = (
+ Path("/usr/share/libdrm/amdgpu.ids"),
+ Path("/usr/local/share/libdrm/amdgpu.ids"),
+ Path("/opt/rocm/share/libdrm/amdgpu.ids"),
+ Path("/opt/amdgpu/share/libdrm/amdgpu.ids"),
+)
+
+
+def _log_rocm_detection(message):
+ if UNSLOTH_ENABLE_LOGGING:
+ logger.info(message)
+
+
+@functools.lru_cache(1)
+def _is_rocm_torch_build() -> bool:
+ # Most official ROCm wheels include a local version suffix like +rocmX.Y.
+ # Some custom/source builds do not, so we fall back to runtime hints.
+ try:
+ torch_version_raw = str(importlib_version("torch")).lower()
+ if "rocm" in torch_version_raw:
+ _log_rocm_detection(
+ "Unsloth: ROCm detection matched torch version tag (+rocm)."
+ )
+ return True
+ except Exception:
+ pass
+
+ # Environment hints commonly present on ROCm runtimes.
+ for key in _ROCM_ENV_HINT_KEYS:
+ value = os.environ.get(key, "")
+ if isinstance(value, str) and value.strip():
+ _log_rocm_detection(
+ f"Unsloth: ROCm detection matched environment key `{key}`."
+ )
+ return True
+
+ # Filesystem / driver hints for ROCm stacks.
+ for path in _ROCM_PATH_HINTS:
+ try:
+ if path.exists():
+ _log_rocm_detection(
+ f"Unsloth: ROCm detection matched filesystem hint `{path}`."
+ )
+ return True
+ except Exception:
+ continue
+
+ _log_rocm_detection("Unsloth: ROCm detection did not match any known hints.")
+ return False
+
+
+def _iter_amdgpu_asic_id_table_candidates():
+ # Try torch-adjacent ids table paths first without importing torch.
+ try:
+ torch_spec = importlib.util.find_spec("torch")
+ except Exception:
+ torch_spec = None
+
+ roots = []
+ if torch_spec is not None:
+ if torch_spec.origin:
+ roots.append(Path(torch_spec.origin).resolve().parent)
+ if torch_spec.submodule_search_locations:
+ for location in torch_spec.submodule_search_locations:
+ roots.append(Path(location).resolve())
+
+ seen = set()
+ for root in roots:
+ for candidate in (
+ root / "share" / "libdrm" / "amdgpu.ids",
+ root.parent / "share" / "libdrm" / "amdgpu.ids",
+ root.parent.parent / "share" / "libdrm" / "amdgpu.ids",
+ ):
+ candidate_str = str(candidate)
+ if candidate_str in seen:
+ continue
+ seen.add(candidate_str)
+ yield candidate
+
+ for candidate in _AMDGPU_ASIC_ID_CANDIDATE_PATHS:
+ candidate_str = str(candidate)
+ if candidate_str in seen:
+ continue
+ seen.add(candidate_str)
+ yield candidate
+
+
+def configure_amdgpu_asic_id_table_path():
+ # Honor an existing valid user-provided path.
+ configured = os.environ.get(_AMDGPU_ASIC_ID_TABLE_PATH_ENV, "").strip()
+ if configured:
+ configured_path = Path(configured)
+ try:
+ if configured_path.is_file():
+ return str(configured_path)
+ except Exception:
+ pass
+
+ # Only attempt this on ROCm-like environments.
+ if not _is_rocm_torch_build():
+ return None
+
+ for candidate in _iter_amdgpu_asic_id_table_candidates():
+ try:
+ if candidate.is_file():
+ os.environ[_AMDGPU_ASIC_ID_TABLE_PATH_ENV] = str(candidate)
+ if UNSLOTH_ENABLE_LOGGING:
+ logger.info(
+ f"Unsloth: Set {_AMDGPU_ASIC_ID_TABLE_PATH_ENV}={candidate}"
+ )
+ return str(candidate)
+ except Exception:
+ continue
+
+ return None
+
+
+def _is_causal_conv1d_name(module_name: str) -> bool:
+ return module_name == _CAUSAL_CONV1D_PREFIX or module_name.startswith(
+ _CAUSAL_CONV1D_PREFIX + "."
+ )
+
+
+def _is_vllm_name(module_name: str) -> bool:
+ return module_name == _VLLM_PREFIX or module_name.startswith(_VLLM_PREFIX + ".")
+
+
+def _resolve_module_name(module_name, package):
+ if not isinstance(module_name, str):
+ return module_name
+ if module_name.startswith("."):
+ try:
+ return importlib.util.resolve_name(module_name, package)
+ except Exception:
+ return module_name
+ return module_name
+
+
+def _is_broken_causal_conv1d_error(error) -> bool:
+ checked = set()
+ current = error
+ while current is not None and id(current) not in checked:
+ checked.add(id(current))
+ message = str(current).lower()
+ if (
+ ("causal_conv1d_cuda" in message and "undefined symbol" in message)
+ or ("_zn3c103hip28c10_hip_check_implementation" in message)
+ or ("causal_conv1d" in message and "undefined symbol" in message)
+ ):
+ return True
+ current = getattr(current, "__cause__", None) or getattr(
+ current, "__context__", None
+ )
+ return False
+
+
+def _is_broken_vllm_error(error) -> bool:
+ checked = set()
+ current = error
+ while current is not None and id(current) not in checked:
+ checked.add(id(current))
+ message = str(current).lower()
+ if (
+ ("vllm/_c" in message or "vllm._c" in message)
+ and (
+ "undefined symbol" in message
+ or "cannot open shared object file" in message
+ or ".so:" in message
+ )
+ ) or ("vllm" in message and "undefined symbol" in message):
+ return True
+ # Also catch CUDA shared library mismatches during vllm import
+ # e.g. "libcudart.so.12: cannot open shared object file"
+ if (
+ "libcudart" in message or "libcublas" in message or "libnvrtc" in message
+ ) and "cannot open shared object file" in message:
+ return True
+ current = getattr(current, "__cause__", None) or getattr(
+ current, "__context__", None
+ )
+ return False
+
+
+def _get_vllm_cuda_mismatch_message(error):
+ """If the error is a CUDA version mismatch, return a helpful install message."""
+ import re as _re
+
+ checked = set()
+ current = error
+ wanted_cuda = None
+ while current is not None and id(current) not in checked:
+ checked.add(id(current))
+ message = str(current)
+ # Extract the CUDA version vllm was built for, e.g. "libcudart.so.12"
+ match = _re.search(r"libcudart\.so\.(\d+)", message)
+ if match:
+ wanted_cuda = match.group(1)
+ break
+ current = getattr(current, "__cause__", None) or getattr(
+ current, "__context__", None
+ )
+ if wanted_cuda is None:
+ return None
+
+ # Detect what CUDA version is actually available on the system
+ system_cuda_display = None # Human-readable, e.g. "13.0"
+ system_cuda_tag = None # For wheel URL, e.g. "130"
+ try:
+ import torch
+
+ cuda_version = torch.version.cuda # e.g. "13.0" or "12.8"
+ if cuda_version:
+ system_cuda_display = cuda_version
+ system_cuda_tag = cuda_version.replace(".", "")[:3] # "130" or "128"
+ except Exception:
+ pass
+
+ if system_cuda_tag is None or system_cuda_tag.startswith(wanted_cuda):
+ return None # Not a mismatch or can't determine
+
+ try:
+ vllm_version = importlib_version("vllm").split("+")[0]
+ except Exception:
+ vllm_version = "VLLM_VERSION"
+
+ cpu_arch = "x86_64"
+ try:
+ import platform
+
+ cpu_arch = platform.machine()
+ except Exception:
+ pass
+
+ return (
+ f"Unsloth: vLLM was built for CUDA {wanted_cuda} but this system has "
+ f"CUDA {system_cuda_display}. Please reinstall vLLM with the correct CUDA version:\n"
+ f"\n"
+ f" uv pip install https://github.com/vllm-project/vllm/releases/download/"
+ f"v{vllm_version}/vllm-{vllm_version}+cu{system_cuda_tag}-cp38-abi3-"
+ f"manylinux_2_35_{cpu_arch}.whl"
+ )
+
+
+class _CausalConv1dImportBlockerLoader(importlib.abc.Loader):
+ __slots__ = ("module_name",)
+
+ def __init__(self, module_name):
+ self.module_name = module_name
+
+ def create_module(self, spec):
+ return None
+
+ def exec_module(self, module):
+ raise ModuleNotFoundError(f"No module named '{self.module_name}'")
+
+
+class _CausalConv1dImportBlockerFinder(importlib.abc.MetaPathFinder):
+ __slots__ = (_CAUSAL_CONV1D_BLOCKER_SENTINEL,)
+
+ def __init__(self):
+ setattr(self, _CAUSAL_CONV1D_BLOCKER_SENTINEL, True)
+
+ def find_spec(self, fullname, path = None, target = None):
+ if not CAUSAL_CONV1D_BROKEN or not _is_causal_conv1d_name(fullname):
+ return None
+ return importlib.machinery.ModuleSpec(
+ name = fullname,
+ loader = _CausalConv1dImportBlockerLoader(fullname),
+ is_package = fullname == _CAUSAL_CONV1D_PREFIX,
+ )
+
+
+class _VllmImportBlockerLoader(importlib.abc.Loader):
+ __slots__ = ("module_name",)
+
+ def __init__(self, module_name):
+ self.module_name = module_name
+
+ def create_module(self, spec):
+ return None
+
+ def exec_module(self, module):
+ raise ModuleNotFoundError(f"No module named '{self.module_name}'")
+
+
+class _VllmImportBlockerFinder(importlib.abc.MetaPathFinder):
+ __slots__ = (_VLLM_BLOCKER_SENTINEL,)
+
+ def __init__(self):
+ setattr(self, _VLLM_BLOCKER_SENTINEL, True)
+
+ def find_spec(self, fullname, path = None, target = None):
+ if not VLLM_BROKEN or not _is_vllm_name(fullname):
+ return None
+ return importlib.machinery.ModuleSpec(
+ name = fullname,
+ loader = _VllmImportBlockerLoader(fullname),
+ is_package = fullname == _VLLM_PREFIX,
+ )
+
+
+def _patch_find_spec_for_causal_conv1d():
+ current_find_spec = importlib.util.find_spec
+ if getattr(current_find_spec, "_unsloth_causal_conv1d_find_spec_patch", False):
+ return
+
+ def _blocked_find_spec(name, package = None):
+ resolved_name = _resolve_module_name(name, package)
+ if CAUSAL_CONV1D_BROKEN and isinstance(resolved_name, str):
+ if _is_causal_conv1d_name(resolved_name):
+ return None
+ return current_find_spec(name, package)
+
+ _blocked_find_spec._unsloth_causal_conv1d_find_spec_patch = True
+ _blocked_find_spec._unsloth_original_find_spec = current_find_spec
+ importlib.util.find_spec = _blocked_find_spec
+
+
+def _patch_find_spec_for_vllm():
+ current_find_spec = importlib.util.find_spec
+ if getattr(current_find_spec, "_unsloth_vllm_find_spec_patch", False):
+ return
+
+ def _blocked_find_spec(name, package = None):
+ resolved_name = _resolve_module_name(name, package)
+ if VLLM_BROKEN and isinstance(resolved_name, str):
+ if _is_vllm_name(resolved_name):
+ return None
+ return current_find_spec(name, package)
+
+ _blocked_find_spec._unsloth_vllm_find_spec_patch = True
+ _blocked_find_spec._unsloth_original_find_spec = current_find_spec
+ importlib.util.find_spec = _blocked_find_spec
+
+
+def _install_causal_conv1d_blocker():
+ _patch_find_spec_for_causal_conv1d()
+ for finder in sys.meta_path:
+ if getattr(finder, _CAUSAL_CONV1D_BLOCKER_SENTINEL, False):
+ return
+ sys.meta_path.insert(0, _CausalConv1dImportBlockerFinder())
+
+
+def _install_vllm_blocker():
+ _patch_find_spec_for_vllm()
+ for finder in sys.meta_path:
+ if getattr(finder, _VLLM_BLOCKER_SENTINEL, False):
+ return
+ sys.meta_path.insert(0, _VllmImportBlockerFinder())
+
+
+def _clear_causal_conv1d_modules():
+ for module_name in list(sys.modules):
+ if _is_causal_conv1d_name(module_name):
+ sys.modules.pop(module_name, None)
+
+
+def _clear_vllm_modules():
+ for module_name in list(sys.modules):
+ if _is_vllm_name(module_name):
+ sys.modules.pop(module_name, None)
+
+
+def disable_broken_vllm(error = None):
+ """Disable vLLM dynamically when its shared library is ABI-broken."""
+ global VLLM_BROKEN
+ if VLLM_BROKEN:
+ _install_vllm_blocker()
+ return True
+
+ failure = error
+ if failure is None:
+ try:
+ if importlib.util.find_spec("vllm") is None:
+ return False
+ except Exception:
+ return False
+
+ try:
+ import vllm # noqa: F401
+
+ return False
+ except Exception as import_error:
+ failure = import_error
+
+ if not _is_broken_vllm_error(failure):
+ return False
+
+ VLLM_BROKEN = True
+ _clear_vllm_modules()
+ _install_vllm_blocker()
+ cuda_msg = _get_vllm_cuda_mismatch_message(failure)
+ if cuda_msg:
+ logger.warning(cuda_msg)
+ else:
+ logger.warning(
+ "Unsloth: Detected broken vLLM binary extension; "
+ "disabling vLLM imports and continuing import.\n"
+ "Please reinstall via `uv pip install unsloth vllm torchvision torchaudio "
+ "--torch-backend=auto`."
+ )
+ return True
+
+
+def _disable_transformers_causal_conv1d():
+ try:
+ import transformers.utils.import_utils as tf_import_utils
+ except Exception:
+ return
+
+ if hasattr(tf_import_utils, "is_causal_conv1d_available"):
+ tf_import_utils.is_causal_conv1d_available = lambda: False
+
+ for attr_name in (
+ "_causal_conv1d_available",
+ "_is_causal_conv1d_available",
+ ):
+ if hasattr(tf_import_utils, attr_name):
+ setattr(tf_import_utils, attr_name, False)
+
+
+def disable_broken_causal_conv1d():
+ """Disable causal_conv1d dynamically when its shared library is ABI-broken.
+
+ This mirrors Unsloth's FlashAttention fallback behavior: if importing causal_conv1d
+ fails with a known binary symbol error, we disable it at startup so model imports do
+ not hard-fail.
+ """
+ global CAUSAL_CONV1D_BROKEN
+ if CAUSAL_CONV1D_BROKEN:
+ _install_causal_conv1d_blocker()
+ _disable_transformers_causal_conv1d()
+ return
+
+ try:
+ if importlib.util.find_spec("causal_conv1d") is None:
+ return
+ except Exception:
+ return
+
+ try:
+ import causal_conv1d # noqa: F401
+
+ return
+ except Exception as error:
+ if not _is_broken_causal_conv1d_error(error):
+ return
+
+ CAUSAL_CONV1D_BROKEN = True
+ _clear_causal_conv1d_modules()
+ _install_causal_conv1d_blocker()
+ _disable_transformers_causal_conv1d()
+ print(
+ "Unsloth: Detected broken causal_conv1d binary; "
+ "disabling causal_conv1d fast path and continuing import."
+ )
diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py
index fbb14013ff..d92229314f 100644
--- a/unsloth/kernels/cross_entropy_loss.py
+++ b/unsloth/kernels/cross_entropy_loss.py
@@ -301,6 +301,7 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
vocab_size: int
n_rows, vocab_size = logits.shape
device = logits.device
+ labels = labels.to(device)
div, mod = divmod(vocab_size, MAX_FUSED_SIZE)
n_chunks: int = div + (mod != 0)
@@ -437,6 +438,7 @@ def fast_cross_entropy_loss(
batch, seq_len, d = logits.shape
assert labels.shape == (batch, seq_len)
+ device = logits.device
loss = Fast_CrossEntropyLoss.apply(
logits.view(batch * seq_len, d),
labels.view(-1),
@@ -445,6 +447,8 @@ def fast_cross_entropy_loss(
)
if n_items is None:
n_items = torch.count_nonzero(labels != -100)
+ if torch.is_tensor(n_items):
+ n_items = n_items.to(device)
return loss.sum() / n_items
diff --git a/unsloth/kernels/fp8.py b/unsloth/kernels/fp8.py
index e9f9161709..a57f4ffb64 100644
--- a/unsloth/kernels/fp8.py
+++ b/unsloth/kernels/fp8.py
@@ -95,8 +95,11 @@ def weight_dequant_block(
def weight_dequant(x: torch.Tensor, s: torch.Tensor, dtype = torch.bfloat16):
- if s.shape[1] == 1:
- # this is row quantized weight, just simple multiplication suffices
+ # Per-tensor scale: single value for entire weight matrix
+ if s.numel() == 1:
+ return x.to(dtype) * s.view(1, 1).to(dtype)
+ # Row quantized weight: scale shape is (m, 1) or (n, 1)
+ elif s.ndim == 2 and s.shape[1] == 1:
if x.shape[0] == s.shape[0]:
y = x.to(dtype) * s.to(dtype)
elif x.shape[1] == s.shape[0]:
@@ -106,8 +109,8 @@ def weight_dequant(x: torch.Tensor, s: torch.Tensor, dtype = torch.bfloat16):
else:
raise ValueError(f"Incompatible shapes {x.shape = }, {s.shape = }")
return y
+ # Block quantized weight: scale shape is (ceil(m/block_m), ceil(n/block_n))
else:
- # this is block quantized weight
return weight_dequant_block(x, s, dtype = dtype)
@@ -238,44 +241,29 @@ def w8a8_block_fp8_matmul_triton(
block_size: list[int],
output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
- """This function performs matrix multiplication with block-wise
- quantization.
- It takes two input tensors `A` and `B` with scales `As` and `Bs`.
- The output is returned in the specified `output_dtype`.
- Args:
- A: The input tensor, e.g., activation.
- B: The input tensor, e.g., weight.
- As: The per-token-group quantization scale for `A`.
- Bs: The per-block quantization scale for `B`.
- block_size: The block size for per-block quantization. It should
- be 2-dim, e.g., [128, 128].
- output_dytpe: The dtype of the returned tensor.
- Returns:
- torch.Tensor: The result of matmul.
- """
- assert len(block_size) == 2
- block_n, block_k = block_size[0], block_size[1]
+ """Block-wise FP8 matmul."""
+ if block_size is None:
+ block_n, block_k = 128, 128
+ else:
+ assert len(block_size) == 2
+ block_n, block_k = block_size[0], block_size[1]
+ N, K = B.shape
assert A.shape[-1] == B.shape[-1]
assert A.shape[:-1] == As.shape[:-1] and A.is_contiguous()
assert triton.cdiv(A.shape[-1], block_k) == As.shape[-1]
- M = A.numel() // A.shape[-1]
-
assert B.ndim == 2 and B.is_contiguous() and Bs.ndim == 2
- N, K = B.shape
assert triton.cdiv(N, block_n) == Bs.shape[0]
assert triton.cdiv(K, block_k) == Bs.shape[1]
+ M = A.numel() // A.shape[-1]
C_shape = A.shape[:-1] + (N,)
C = A.new_empty(C_shape, dtype = output_dtype)
BLOCK_SIZE_M = 128
if M < BLOCK_SIZE_M:
- BLOCK_SIZE_M = triton.next_power_of_2(M)
- BLOCK_SIZE_M = max(BLOCK_SIZE_M, 16)
- BLOCK_SIZE_K = block_k
- assert block_k % BLOCK_SIZE_K == 0
- BLOCK_SIZE_N = block_n
+ BLOCK_SIZE_M = max(triton.next_power_of_2(M), 16)
+ BLOCK_SIZE_K, BLOCK_SIZE_N = block_k, block_n
def grid(META):
return (
@@ -342,29 +330,41 @@ fp8_block_matmul = (
class FP8BlockQuantLinear(torch.autograd.Function):
@staticmethod
def forward(ctx, X, weight, weight_scale):
- # block_size = getattr(weight, 'block_size', [128,128])
m, n = weight.shape
- p, q = weight_scale.shape
- block_size = getattr(weight, "block_size", None) or getattr(
- weight_scale, "block_size", [128, 128]
- )
- assert block_size is not None, "block_size is not set"
- if triton.cdiv(m, block_size[0]) != p or triton.cdiv(n, block_size[1]) != q:
- if (
- triton.cdiv(m, block_size[0]) == q
- and triton.cdiv(n, block_size[1]) == p
- ):
- # weights are transposed during backward pass for training :)
- # We transpose weight scale to counter that. Note that transposing weight would cause issues with matmul with input X
- weight_scale = weight_scale.T
- else:
- raise ValueError(
- f"Weight shape {weight.shape} and scales shape {weight_scale.shape} is not compatible with block size {block_size}"
- )
+
+ # Save original scale for backward (before any transformation)
+ original_weight_scale = weight_scale
+
+ # Handle per-tensor quantization: expand scalar to block scale shape
+ if weight_scale.numel() == 1:
+ block_size = [128, 128]
+ # Expand scalar to (ceil(m/128), ceil(n/128)) - same value for all blocks
+ num_blocks_m = triton.cdiv(m, block_size[0])
+ num_blocks_n = triton.cdiv(n, block_size[1])
+ weight_scale = weight_scale.expand(num_blocks_m, num_blocks_n).contiguous()
+ else:
+ # Block quantization path
+ p, q = weight_scale.shape
+ block_size = getattr(weight, "block_size", None) or getattr(
+ weight_scale, "block_size", [128, 128]
+ )
+ assert block_size is not None, "block_size is not set"
+ if triton.cdiv(m, block_size[0]) != p or triton.cdiv(n, block_size[1]) != q:
+ if (
+ triton.cdiv(m, block_size[0]) == q
+ and triton.cdiv(n, block_size[1]) == p
+ ):
+ weight_scale = weight_scale.T
+ original_weight_scale = weight_scale # Update for transposed case
+ else:
+ raise ValueError(
+ f"Weight shape {weight.shape} and scales shape {weight_scale.shape} is not compatible with block size {block_size}"
+ )
if not weight.is_contiguous():
weight = weight.contiguous()
- # this is replica of https://github.com/huggingface/transformers/blob/01c9e1ba683b3e50d7c76bf92f2d470759fd5e81/src/transformers/integrations/finegrained_fp8.py#L331-L353
+
+ # Quantize input and run FP8 matmul
qinput, scale = act_quant(X, block_size[1])
output = fp8_block_matmul(
qinput,
@@ -375,8 +375,7 @@ class FP8BlockQuantLinear(torch.autograd.Function):
output_dtype = X.dtype,
)
ctx.weight = weight
- ctx.weight_scale = weight_scale
- ctx.block_size = block_size
+ ctx.weight_scale = original_weight_scale # Save original for backward
return output.to(X.dtype)
@staticmethod
@@ -580,7 +579,14 @@ try:
if Version(fbgemm_gpu.__version__) >= Version("1.4.0"):
# We must manually confirm if blockwise FBGEMM works!
# This check is a must for consumer grade GPUs which fail
- if test_has_fbgemm():
+ # Suppress CUDA device printf during probe -- on Blackwell (SM100) GPUs,
+ # FBGEMM's CUTLASS blockwise kernel (hardcoded SM90) fires thousands of
+ # "Arch conditional MMA" lines to stdout fd 1 before aborting.
+ from unsloth.import_fixes import suppress_cuda_printf
+
+ with suppress_cuda_printf():
+ _has_fbgemm = test_has_fbgemm()
+ if _has_fbgemm:
os.environ["UNSLOTH_HAS_FBGEMM"] = "1"
logger.info(f"Using fbgemm_gpu block quantized FP8 matmul")
fp8_block_quant_linear = fp8_fbgemm_block_linear
@@ -592,11 +598,14 @@ except:
@torch_compile
def fp8_linear(X, weight, weight_scale, bias = None):
- if weight_scale.ndim == 2 and weight_scale.shape[1] > 1:
- # This is block quantized FP8 matmul
+ # Per-tensor quantization: single scalar scale for entire weight
+ # Block quantized FP8: 2D scale tensor with multiple columns
+ if weight_scale.numel() == 1 or (
+ weight_scale.ndim == 2 and weight_scale.shape[1] > 1
+ ):
out = fp8_block_quant_linear(X, weight, weight_scale)
+ # Row/channel quantized FP8: 2D scale with shape (n, 1)
else:
- # Row quantized FP8
out = fbgemm_fp8_linear(X, weight, weight_scale, bias)
return out
diff --git a/unsloth/kernels/layernorm.py b/unsloth/kernels/layernorm.py
index 5e2e3af2f8..9e64c3d341 100644
--- a/unsloth/kernels/layernorm.py
+++ b/unsloth/kernels/layernorm.py
@@ -55,7 +55,9 @@ def layernorm_forward(
# (X[0] - mean) == -mean so we need to mask it out
XX = tl.where(mask, X_row - mean_X, 0)
row_var = tl.sum(XX * XX, axis = 0) / n_cols
- inv_var = tl.math.rsqrt(row_var + eps)
+ # Explicit float32 scalar to ensure correct type promotion on HIP/ROCm
+ eps_f32 = tl.full((), eps, tl.float32)
+ inv_var = tl.math.rsqrt(row_var + eps_f32)
tl.store(r, inv_var)
tl.store(mu, mean_X)
output = (XX * inv_var) * W_row + b_row
diff --git a/unsloth/kernels/moe/grouped_gemm/interface.py b/unsloth/kernels/moe/grouped_gemm/interface.py
index 572b02702b..5588458973 100644
--- a/unsloth/kernels/moe/grouped_gemm/interface.py
+++ b/unsloth/kernels/moe/grouped_gemm/interface.py
@@ -4,6 +4,7 @@
import logging
import warnings
from dataclasses import asdict
+from unsloth import DEVICE_TYPE
import torch
import triton
@@ -38,9 +39,11 @@ logger.addHandler(ch)
# Precompute TMA support to avoid graph breaks
# TMA requires both:
-# 1. GPU capability >= 9 (Hopper+)
+# 1. NVIDIA GPU with capability >= 9 (Hopper+)
# 2. Triton version with TMA API (make_tensor_descriptor or _experimental_make_tensor_descriptor)
def _check_tma_support():
+ if DEVICE_TYPE in ("xpu", "hip"):
+ return False
import triton.language as tl
gpu_supports_tma = torch.cuda.get_device_capability()[0] >= 9
diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py
index 82e0cd0e9b..74c16c1e63 100644
--- a/unsloth/kernels/rms_layernorm.py
+++ b/unsloth/kernels/rms_layernorm.py
@@ -49,7 +49,9 @@ def _rms_layernorm_forward(
W_row = tl.load(W + col_offsets, mask = mask, other = 0) # .to(tl.float32)
row_var = tl.sum(X_row * X_row, axis = 0) / n_cols
- inv_var = tl.math.rsqrt(row_var + eps)
+ # Explicit float32 scalar to ensure correct type promotion on HIP/ROCm
+ eps_f32 = tl.full((), eps, tl.float32)
+ inv_var = tl.math.rsqrt(row_var + eps_f32)
tl.store(r, inv_var)
normed = X_row * inv_var
normed = normed.to(W_row.dtype) # Exact copy from HF
@@ -147,7 +149,9 @@ def _gemma_rms_layernorm_forward(
W_row = tl.load(W + col_offsets, mask = mask, other = 0).to(tl.float32)
row_var = tl.sum(X_row * X_row, axis = 0) / n_cols
- inv_var = tl.math.rsqrt(row_var + eps)
+ # Explicit float32 scalar to ensure correct type promotion on HIP/ROCm
+ eps_f32 = tl.full((), eps, tl.float32)
+ inv_var = tl.math.rsqrt(row_var + eps_f32)
tl.store(r, inv_var)
normed = X_row * inv_var
output = normed * (W_row + 1.0)
diff --git a/unsloth/kernels/utils.py b/unsloth/kernels/utils.py
index 5dcc7c232c..90f2d5d238 100644
--- a/unsloth/kernels/utils.py
+++ b/unsloth/kernels/utils.py
@@ -82,6 +82,18 @@ def is_cdna():
"gfx940",
"gfx941",
"gfx942",
+ "gfx950", # CDNA4 (MI350/MI355X)
+ )
+
+
+@functools.lru_cache(1)
+def is_rdna():
+ """Detect ROCm-supported RDNA consumer/workstation GPUs (RDNA3, RDNA4)."""
+ return is_hip() and triton.runtime.driver.active.get_current_target().arch in (
+ "gfx1100",
+ "gfx1101",
+ "gfx1200",
+ "gfx1201",
)
@@ -388,7 +400,7 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM:
global ABSMAX_BUFFERS
WEIGHT_BUFFER = WEIGHT_BUFFERS[device_index]
ABSMAX_BUFFER = ABSMAX_BUFFERS[device_index]
- if WEIGHT_BUFFER is None:
+ if WEIGHT_BUFFER is None or WEIGHT_BUFFER.dtype != dtype:
WEIGHT_BUFFERS[device_index] = WEIGHT_BUFFER = torch_empty(
size, dtype = dtype, device = device, requires_grad = False
)
@@ -498,7 +510,7 @@ elif DEVICE_TYPE in ("cuda", "hip") and HAS_CUDA_STREAM:
global ABSMAX_BUFFERS
WEIGHT_BUFFER = WEIGHT_BUFFERS[device_index]
ABSMAX_BUFFER = ABSMAX_BUFFERS[device_index]
- if WEIGHT_BUFFER is None:
+ if WEIGHT_BUFFER is None or WEIGHT_BUFFER.dtype != dtype:
WEIGHT_BUFFERS[device_index] = WEIGHT_BUFFER = torch_empty(
size, dtype = dtype, device = device, requires_grad = False
)
diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py
index 70d71b3e06..b74743d724 100644
--- a/unsloth/models/_utils.py
+++ b/unsloth/models/_utils.py
@@ -12,7 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-__version__ = "2026.2.1"
+__version__ = "2026.3.1"
__all__ = [
"SUPPORTS_BFLOAT16",
@@ -28,6 +28,7 @@ __all__ = [
"HAS_FLASH_ATTENTION_SOFTCAPPING",
"USE_MODELSCOPE",
"platform_system",
+ "resolve_hip_gpu_stats_name",
"patch_tokenizer",
"get_statistics",
"Unsloth_Offloaded_Gradient_Checkpointer",
@@ -149,6 +150,39 @@ from unsloth_zoo.compiler import (
from unsloth_zoo.training_utils import (
prepare_model_for_training,
)
+
+
+def resolve_hip_gpu_stats_name(gpu_stats):
+ name = str(getattr(gpu_stats, "name", "") or "").strip()
+ name = re.sub(r"\s*\([^)]*\)\s*$", "", name).strip()
+ normalized_name = name.lower().strip(". ")
+ if normalized_name and normalized_name not in ("amd radeon graphics",):
+ return name + ". "
+
+ try:
+ torch_name = str(torch.cuda.get_device_name(0) or "").strip()
+ torch_name = re.sub(r"\s*\([^)]*\)\s*$", "", torch_name).strip()
+ except Exception:
+ torch_name = ""
+ normalized_torch_name = torch_name.lower().strip(". ")
+ if normalized_torch_name and normalized_torch_name not in ("amd radeon graphics",):
+ return torch_name + ". "
+
+ arch_name = ""
+ for key in ("gcnArchName", "gcn_arch_name", "arch_name", "gfx_arch_name"):
+ value = getattr(gpu_stats, key, None)
+ if value is not None and str(value).strip():
+ arch_name = str(value).strip()
+ break
+
+ if arch_name:
+ arch_name = arch_name.strip()
+ match = re.search(r"(gfx[0-9a-z]+)", arch_name, flags = re.I)
+ if match:
+ return f"AMD {match.group(1).lower()} GPU. "
+ return "AMD GPU. "
+
+
from unsloth_zoo.temporary_patches import (
TEMPORARY_PATCHES,
)
@@ -200,11 +234,16 @@ def prefer_flex_attn_if_supported(model_class, config):
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.
+ # GPT-OSS, Mllama and Gemma3N use eager/sdpa attention during
+ # inference since flex attention returns incorrect results or errors out.
+ # GPT-OSS: left padding issues cause incorrect outputs.
+ # Mllama: _update_causal_mask uses make_flex_block_causal_mask which
+ # creates BlockMask with Q_LEN=KV_LEN=total_seq_len, but during
+ # decode q_len=1, causing ValueError. Needs transformers update.
+ # Gemma3N: timm vision wrappers (eg Gemma3nVisionConfig) do not
+ # support flex_attention.
model_type = getattr(config, "model_type", "") if config else ""
- if model_type == "gpt_oss":
+ if model_type in ("gpt_oss", "mllama") or str(model_type).startswith("gemma3n"):
return None
if config is not None:
setattr(config, "_attn_implementation", "flex_attention")
@@ -842,11 +881,8 @@ if DEVICE_TYPE == "cuda":
)
except:
print(
- "Unsloth: Your Flash Attention 2 installation seems to be broken?\n"
- "A possible explanation is you have a new CUDA version which isn't\n"
- "yet compatible with FA2? Please file a ticket to Unsloth or FA2.\n"
- "We shall now use Xformers instead, which does not have any performance hits!\n"
- "We found this negligible impact by benchmarking on 1x A100."
+ "Unsloth: Your Flash Attention 2 installation seems to be broken. "
+ "Using Xformers instead. No performance changes will be seen."
)
# Stop Flash Attention from importing!
@@ -894,11 +930,8 @@ elif DEVICE_TYPE == "hip":
)
except:
print(
- "Unsloth: Your Flash Attention 2 installation seems to be broken?\n"
- "A possible explanation is you have a new CUDA version which isn't\n"
- "yet compatible with FA2? Please file a ticket to Unsloth or FA2.\n"
- "We shall now use Xformers instead, which does not have any performance hits!\n"
- "We found this negligible impact by benchmarking on 1x A100."
+ "Unsloth: Your Flash Attention 2 installation seems to be broken. "
+ "Using Xformers instead. No performance changes will be seen."
)
# Stop Flash Attention from importing!
@@ -1439,7 +1472,7 @@ BitsAndBytesConfig__init__ = BitsAndBytesConfig__init__.replace(
)
exec(BitsAndBytesConfig__init__, globals())
-if DEVICE_COUNT == 1:
+if DEVICE_COUNT == 1 and int(os.environ.get("WORLD_SIZE", "1")) <= 1:
from accelerate.utils.dataclasses import DistributedType
def _prepare_backend(self, *args, **kwargs):
@@ -2538,6 +2571,14 @@ def patch_hf_quantizer():
except Exception as e:
logger.warning(f"Failed to patch FbgemmFp8HfQuantizer. Error {e}")
+ try:
+ from transformers.quantizers.quantizer_torchao import TorchAoHfQuantizer
+
+ TorchAoHfQuantizer.is_trainable = property(make_trainable)
+ TorchAoHfQuantizer.is_qat_trainable = property(make_trainable)
+ except Exception as e:
+ logger.warning(f"Failed to patch TorchAoHfQuantizer. Error {e}")
+
patch_hf_quantizer()
diff --git a/unsloth/models/cohere.py b/unsloth/models/cohere.py
index c33317ee02..4251f3acd9 100644
--- a/unsloth/models/cohere.py
+++ b/unsloth/models/cohere.py
@@ -188,7 +188,9 @@ def CohereDecoderLayer_fast_forward(
self, "_flag_for_generation"
): # past_key_value is not None:
out_weight = torch.empty(
- self.input_layernorm.weight.shape, dtype = torch.float32, device = "cuda:0"
+ self.input_layernorm.weight.shape,
+ dtype = torch.float32,
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
# Self Attention
@@ -254,6 +256,7 @@ def CohereAttention_fast_forward_inference(
position_ids,
do_prefill = False,
attention_mask = None,
+ **kwargs,
):
Xn = hidden_states
bsz, _, hd = hidden_states.size()
@@ -277,26 +280,28 @@ def CohereAttention_fast_forward_inference(
self.paged_attention = torch.empty(
(KV_CACHE_INCREMENT + seq_len + 1, 2, bsz, n_kv_heads, head_dim),
dtype = dtype,
- device = "cuda:0",
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
self.paged_attention_K = self.paged_attention[:, 0]
self.paged_attention_V = self.paged_attention[:, 1]
self.paged_attention_K[:seq_len] = K1.permute(2, 0, 1, 3)
self.paged_attention_V[:seq_len] = V1.permute(2, 0, 1, 3)
self.temp_QA = torch.empty(
- (2, bsz, 1, attention_size), dtype = dtype, device = "cuda:0"
+ (2, bsz, 1, attention_size), dtype = dtype, device = f"{DEVICE_TYPE_TORCH}:0"
)
self.temp_KV = torch.empty(
- (2, bsz, 1, n_kv_heads * head_dim), dtype = dtype, device = "cuda:0"
+ (2, bsz, 1, n_kv_heads * head_dim),
+ dtype = dtype,
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
self.RH_Q = torch.empty(
- (bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda:0"
+ (bsz, n_heads, 1, head_dim), dtype = dtype, device = f"{DEVICE_TYPE_TORCH}:0"
)
# Mistral Nemo 12b has weird dimensions
if attention_size != hidden_size:
self.temp_O = torch.empty(
- (1, bsz, hidden_size), dtype = dtype, device = "cuda:0"
+ (bsz, 1, hidden_size), dtype = dtype, device = f"{DEVICE_TYPE_TORCH}:0"
)
else:
self.temp_O = self.temp_QA[1][:, :, :hidden_size]
@@ -304,17 +309,21 @@ def CohereAttention_fast_forward_inference(
self.attention = torch.empty(
(bsz, n_heads, 1, KV_CACHE_INCREMENT + seq_len),
dtype = dtype,
- device = "cuda:0",
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
self.scalar = 1.0 / math_sqrt(self.head_dim)
self.half_head_dim = head_dim // 2
# Cohere has QK layernorms
if self.use_qk_norm:
self.q_norm_out_weight = torch.empty(
- self.q_norm.weight.shape, dtype = torch.float32, device = "cuda:0"
+ self.q_norm.weight.shape,
+ dtype = torch.float32,
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
self.k_norm_out_weight = torch.empty(
- self.k_norm.weight.shape, dtype = torch.float32, device = "cuda:0"
+ self.k_norm.weight.shape,
+ dtype = torch.float32,
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
else:
self.q_norm_out_weight = None
@@ -355,7 +364,7 @@ def CohereAttention_fast_forward_inference(
RH_Q = self.RH_Q
RH_Q[:, :, :, :h] = Qn[:, :, :, h:]
RH_Q[:, :, :, h:] = Qn[:, :, :, :h]
- torch.neg(RH_Q[:, :, :, :h], out = RH_Q[:, :, :, :h])
+ RH_Q[:, :, :, :h].neg_()
Qn *= cos
Qn.addcmul_(RH_Q, sin)
@@ -364,7 +373,7 @@ def CohereAttention_fast_forward_inference(
] # torch.empty((n_kv_heads, 1, head_dim), dtype = dtype, device = "cuda:0")
RH_K[:, :, :, :h] = Kn[:, :, :, h:]
RH_K[:, :, :, h:] = Kn[:, :, :, :h]
- torch.neg(RH_K[:, :, :, :h], out = RH_K[:, :, :, :h])
+ RH_K[:, :, :, :h].neg_()
Kn *= cos
Kn.addcmul_(RH_K, sin)
@@ -379,10 +388,11 @@ def CohereAttention_fast_forward_inference(
# Handle sliding windows
sliding_window = getattr(self.config, "sliding_window", None)
if sliding_window is not None and kv_seq_len > sliding_window:
- # From https://github.com/huggingface/transformers/blob/main/src/transformers/models/mistral/modeling_mistral.py#L193
- slicing_tokens = 1 - sliding_window
- Knn = Kn[:, :, slicing_tokens:, :] # .contiguous()
- Vnn = Vn[:, :, slicing_tokens:, :] # .contiguous()
+ start = kv_seq_len - sliding_window
+ Knn = Kn[:, :, start:, :] # .contiguous()
+ Vnn = Vn[:, :, start:, :] # .contiguous()
+ if attention_mask is not None:
+ attention_mask = attention_mask[..., start:]
else:
Knn, Vnn = Kn, Vn
@@ -397,9 +407,6 @@ def CohereAttention_fast_forward_inference(
)
Knn = Knn.reshape(bsz, n_heads, cached_len, head_dim)
Vnn = Vnn.reshape(bsz, n_heads, cached_len, head_dim)
- # else:
- # Knn, Vnn = Knn, Vnn
- # pass
# Attention
if bsz == 1:
@@ -408,7 +415,6 @@ def CohereAttention_fast_forward_inference(
A = torch_matmul(
Qn, Knn.transpose(2, 3), out = self.attention[:, :, :, :cached_len]
)
- # if attention_mask is not None: A += attention_mask # Must add attention_mask for batched
A[:] = torch_nn_functional_softmax(
A, dim = -1, dtype = torch.float32
) # .to(A.dtype)
@@ -453,6 +459,9 @@ def CohereModel_fast_forward_inference(
seq_len,
sliding_window = getattr(self.config, "sliding_window", None),
)
+ # Pre-convert to bool once for all layers (avoids per-layer .eq(0))
+ if attention_mask is not None and attention_mask.dtype != torch.bool:
+ attention_mask = attention_mask.eq(0)
else:
attention_mask = None
diff --git a/unsloth/models/falcon_h1.py b/unsloth/models/falcon_h1.py
index 428f49d727..6e3b16b21b 100644
--- a/unsloth/models/falcon_h1.py
+++ b/unsloth/models/falcon_h1.py
@@ -197,6 +197,7 @@ def FalconH1Attention_fast_forward_inference(
position_ids,
do_prefill = False,
attention_mask = None,
+ **kwargs,
):
"""
https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L406
@@ -265,7 +266,7 @@ def FalconH1Attention_fast_forward_inference(
# Mistral Nemo 12b has weird dimensions
if attention_size != hidden_size:
- self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = device)
+ self.temp_O = torch.empty((bsz, 1, hidden_size), dtype = dtype, device = device)
else:
self.temp_O = self.temp_QA[1][:, :, :hidden_size]
@@ -292,7 +293,7 @@ def FalconH1Attention_fast_forward_inference(
Qn = fast_linear_forward(self.q_proj, Xn, out = self.temp_QA[0])
Kn = fast_linear_forward(self.k_proj, Xn, out = self.temp_KV[0])
- Kn = Kn * self.config.key_multiplier
+ Kn.mul_(self.config.key_multiplier)
Vn = fast_linear_forward(self.v_proj, Xn, out = self.temp_KV[1])
Qn = Qn.view(
bsz, 1, n_heads, head_dim
@@ -343,10 +344,11 @@ def FalconH1Attention_fast_forward_inference(
# Handle sliding windows
sliding_window = getattr(self.config, "sliding_window", None)
if sliding_window is not None and kv_seq_len > sliding_window:
- # From https://github.com/huggingface/transformers/blob/main/src/transformers/models/mistral/modeling_mistral.py#L193
- slicing_tokens = 1 - sliding_window
- Knn = Kn[:, :, slicing_tokens:, :] # .contiguous()
- Vnn = Vn[:, :, slicing_tokens:, :] # .contiguous()
+ start = kv_seq_len - sliding_window
+ Knn = Kn[:, :, start:, :] # .contiguous()
+ Vnn = Vn[:, :, start:, :] # .contiguous()
+ if attention_mask is not None:
+ attention_mask = attention_mask[..., start:]
else:
Knn, Vnn = Kn, Vn
@@ -361,9 +363,6 @@ def FalconH1Attention_fast_forward_inference(
)
Knn = Knn.reshape(bsz, n_heads, cached_len, head_dim)
Vnn = Vnn.reshape(bsz, n_heads, cached_len, head_dim)
- # else:
- # Knn, Vnn = Knn, Vnn
- # pass
# Attention
if bsz == 1:
@@ -372,7 +371,6 @@ def FalconH1Attention_fast_forward_inference(
A = torch_matmul(
Qn, Knn.transpose(2, 3), out = self.attention[:, :, :, :cached_len]
)
- # if attention_mask is not None: A += attention_mask # Must add attention_mask for batched
A[:] = torch_nn_functional_softmax(
A, dim = -1, dtype = torch.float32
) # .to(A.dtype)
@@ -533,11 +531,19 @@ def _FalconH1_fast_forward_inference(
bsz, q_len, hd = X.shape
assert q_len == 1
# Get saved buffers to reduce memory movement
- residual = torch.empty((bsz, q_len, hd), dtype = torch.float32, device = "cuda:0")
- _XX = torch.empty((2, bsz, q_len, hd), dtype = torch.float32, device = "cuda:0")
+ residual = torch.empty(
+ (bsz, q_len, hd), dtype = torch.float32, device = f"{DEVICE_TYPE_TORCH}:0"
+ )
+ _XX = torch.empty(
+ (2, bsz, q_len, hd), dtype = torch.float32, device = f"{DEVICE_TYPE_TORCH}:0"
+ )
XX, XX2 = _XX[0], _XX[1]
- variance = torch.empty((bsz, q_len, 1), dtype = torch.float32, device = "cuda:0")
- temp_mlp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda:0")
+ variance = torch.empty(
+ (bsz, q_len, 1), dtype = torch.float32, device = f"{DEVICE_TYPE_TORCH}:0"
+ )
+ temp_mlp = torch.empty(
+ (2, bsz, 1, mlp_size), dtype = X.dtype, device = f"{DEVICE_TYPE_TORCH}:0"
+ )
temp_gate, temp_up = temp_mlp[0], temp_mlp[1]
seq_len = past_key_values[0][0].shape[-2]
if bsz != 1:
diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py
index 55a8c8697f..cf543ae094 100644
--- a/unsloth/models/gemma.py
+++ b/unsloth/models/gemma.py
@@ -97,7 +97,9 @@ def GemmaDecoderLayer_fast_forward(
self, "_flag_for_generation"
): # past_key_value is not None:
out_weight = torch.empty(
- self.input_layernorm.weight.shape, dtype = torch.float32, device = "cuda:0"
+ self.input_layernorm.weight.shape,
+ dtype = torch.float32,
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
# Self Attention
@@ -191,6 +193,7 @@ def GemmaModel_fast_forward_inference(
bsz, q_len, hd = hidden_states.shape
seq_len = past_key_values[0][0].shape[-2]
+ kv_seq_len = seq_len + 1
if bsz != 1:
attention_mask = _prepare_4d_causal_attention_mask_for_sdpa(
attention_mask,
@@ -198,6 +201,12 @@ def GemmaModel_fast_forward_inference(
hidden_states,
seq_len,
)
+ # Pre-convert to bool once for all layers (avoids per-layer .eq(0))
+ if attention_mask is not None and attention_mask.dtype != torch.bool:
+ attention_mask = attention_mask.eq(0)
+
+ # Compute rotary_seq_len once to avoid per-layer GPU-CPU sync from .item()
+ rotary_seq_len = max(kv_seq_len, int(position_ids.max().item()) + 1)
next_decoder_cache = []
for idx, decoder_layer in enumerate(self.model.layers):
@@ -217,6 +226,7 @@ def GemmaModel_fast_forward_inference(
position_ids = position_ids,
attention_mask = attention_mask,
do_prefill = not hasattr(decoder_layer.self_attn, "paged_attention"),
+ rotary_seq_len = rotary_seq_len,
)
hidden_states += residual
diff --git a/unsloth/models/gemma2.py b/unsloth/models/gemma2.py
index 03e77f6504..e59b8d5ebd 100644
--- a/unsloth/models/gemma2.py
+++ b/unsloth/models/gemma2.py
@@ -222,7 +222,9 @@ def Gemma2DecoderLayer_fast_forward(
self, "_flag_for_generation"
): # past_key_value is not None:
out_weight = torch.empty(
- self.input_layernorm.weight.shape, dtype = torch.float32, device = "cuda:0"
+ self.input_layernorm.weight.shape,
+ dtype = torch.float32,
+ device = f"{DEVICE_TYPE_TORCH}:0",
)
# Self Attention
@@ -352,7 +354,7 @@ def Gemma2Attention_fast_forward_inference(
)
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = device)
# Only for Gemma2
- self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = device)
+ self.temp_O = torch.empty((bsz, 1, hidden_size), dtype = dtype, device = device)
self.attention = torch.empty(
(bsz, n_heads, 1, KV_CACHE_INCREMENT + seq_len), dtype = dtype, device = device
)
@@ -399,7 +401,7 @@ def Gemma2Attention_fast_forward_inference(
RH_Q = self.RH_Q
RH_Q[:, :, :, :h] = Qn[:, :, :, h:]
RH_Q[:, :, :, h:] = Qn[:, :, :, :h]
- torch.neg(RH_Q[:, :, :, :h], out = RH_Q[:, :, :, :h])
+ RH_Q[:, :, :, :h].neg_()
Qn *= cos
Qn.addcmul_(RH_Q, sin)
@@ -408,7 +410,7 @@ def Gemma2Attention_fast_forward_inference(
] # torch.empty((n_kv_heads, 1, head_dim), dtype = dtype, device = "cuda:0")
RH_K[:, :, :, :h] = Kn[:, :, :, h:]
RH_K[:, :, :, h:] = Kn[:, :, :, :h]
- torch.neg(RH_K[:, :, :, :h], out = RH_K[:, :, :, :h])
+ RH_K[:, :, :, :h].neg_()
Kn *= cos
Kn.addcmul_(RH_K, sin)
@@ -423,10 +425,9 @@ def Gemma2Attention_fast_forward_inference(
# Handle sliding windows
sliding_window = self.config.sliding_window
if use_sliding_window and kv_seq_len > sliding_window:
- # From https://github.com/huggingface/transformers/blob/main/src/transformers/models/mistral/modeling_mistral.py#L193
- slicing_tokens = 1 - sliding_window
- Knn = Kn[:, :, slicing_tokens:, :] # .contiguous()
- Vnn = Vn[:, :, slicing_tokens:, :] # .contiguous()
+ start = kv_seq_len - sliding_window
+ Knn = Kn[:, :, start:, :] # .contiguous()
+ Vnn = Vn[:, :, start:, :] # .contiguous()
else:
Knn, Vnn = Kn, Vn
@@ -441,28 +442,32 @@ def Gemma2Attention_fast_forward_inference(
)
Knn = Knn.reshape(bsz, n_heads, cached_len, head_dim)
Vnn = Vnn.reshape(bsz, n_heads, cached_len, head_dim)
- # else:
- # Knn, Vnn = Knn, Vnn
- # pass
# Attention
- # if bsz == 1:
+ # [TODO] Gemma2 uses manual matmul for all batch sizes because SDPA does
+ # not support softcapping (tanh logit scaling). If a future PyTorch adds
+ # a softcap param to scaled_dot_product_attention, consider using SDPA
+ # for bsz > 1 to match the llama/qwen3 pattern.
Qn *= (
self.scalar
) # See https://github.com/ggerganov/llama.cpp/issues/7805#issuecomment-2153349963
# It seems like doing (Q * scalar) @ K is better than (Q @ K) * scalar to stop overflows
A = torch_matmul(Qn, Knn.transpose(2, 3), out = self.attention[:, :, :, :cached_len])
- # if attention_mask is not None: A += attention_mask # Must add attention_mask for batched
+ # Softcapping must happen BEFORE the mask is applied.
+ # Reference: google-deepmind/gemma _modules.py and transformers gemma2 eager_attention_forward
A *= self.reciprocal_t
- torch_tanh(A, out = A)
+ A.tanh_()
A *= self.t # Logit softcapping
+ if attention_mask is not None and isinstance(attention_mask, torch.Tensor):
+ # Slice mask to match K/V when sliding window is active
+ if attention_mask.shape[-1] != A.shape[-1]:
+ attention_mask = attention_mask[:, :, :, -A.shape[-1] :]
+ A += attention_mask
+
A[:] = torch_nn_functional_softmax(A, dim = -1, dtype = torch.float32) # .to(A.dtype)
A = torch_matmul(A, Vnn, out = Qn)
- # else:
- # A = scaled_dot_product_attention(Qn, Knn, Vnn, attn_mask = attention_mask, is_causal = False)
- # pass
A = A.transpose(1, 2)
A = A.reshape(bsz, 1, attention_size)
A = fast_linear_forward(self.o_proj, A, out = self.temp_O)
diff --git a/unsloth/models/granite.py b/unsloth/models/granite.py
index 168df90f4c..79ac41c43f 100644
--- a/unsloth/models/granite.py
+++ b/unsloth/models/granite.py
@@ -323,8 +323,7 @@ def GraniteAttention_fast_forward_inference(
(2, bsz, 1, n_kv_heads * head_dim), dtype = dtype, device = device
)
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = device)
- # Only for Gemma2
- self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = device)
+ self.temp_O = torch.empty((bsz, 1, hidden_size), dtype = dtype, device = device)
self.attention = torch.empty(
(bsz, n_heads, 1, KV_CACHE_INCREMENT + seq_len), dtype = dtype, device = device
)
@@ -362,7 +361,7 @@ def GraniteAttention_fast_forward_inference(
RH_Q = self.RH_Q
RH_Q[:, :, :, :h] = Qn[:, :, :, h:]
RH_Q[:, :, :, h:] = Qn[:, :, :, :h]
- torch.neg(RH_Q[:, :, :, :h], out = RH_Q[:, :, :, :h])
+ RH_Q[:, :, :, :h].neg_()
Qn *= cos
Qn.addcmul_(RH_Q, sin)
@@ -371,7 +370,7 @@ def GraniteAttention_fast_forward_inference(
] # torch.empty((n_kv_heads, 1, head_dim), dtype = dtype, device = "cuda:0")
RH_K[:, :, :, :h] = Kn[:, :, :, h:]
RH_K[:, :, :, h:] = Kn[:, :, :, :h]
- torch.neg(RH_K[:, :, :, :h], out = RH_K[:, :, :, :h])
+ RH_K[:, :, :, :h].neg_()
Kn *= cos
Kn.addcmul_(RH_K, sin)
@@ -385,7 +384,7 @@ def GraniteAttention_fast_forward_inference(
# Grouped query attention
_, _, cached_len, _ = Kn.shape
- if n_groups != 1:
+ if bsz == 1 or ((not SDPA_HAS_GQA) and n_groups != 1):
Kn = Kn[:, :, None, :, :].expand(
bsz, n_kv_heads, n_groups, cached_len, head_dim
)
@@ -394,20 +393,39 @@ def GraniteAttention_fast_forward_inference(
)
Kn = Kn.reshape(bsz, n_heads, cached_len, head_dim)
Vn = Vn.reshape(bsz, n_heads, cached_len, head_dim)
- # else:
- # Kn, Vn = Kn, Vn
- # pass
- Qn *= self.scaling
- A = torch_matmul(Qn, Kn.transpose(2, 3), out = self.attention[:, :, :, :cached_len])
-
- # if attention_mask is not None: A += attention_mask # Must add attention_mask for batched
-
- A[:] = torch_nn_functional_softmax(A, dim = -1, dtype = torch.float32) # .to(A.dtype)
- A = torch_matmul(A, Vn, out = Qn)
- # else:
- # A = scaled_dot_product_attention(Qn, Kn, Vn, attn_mask = attention_mask, is_causal = False)
- # pass
+ # Attention
+ if bsz == 1:
+ Qn *= self.scaling
+ A = torch_matmul(
+ Qn, Kn.transpose(2, 3), out = self.attention[:, :, :, :cached_len]
+ )
+ A[:] = torch_nn_functional_softmax(A, dim = -1, dtype = torch.float32)
+ A = torch_matmul(A, Vn, out = Qn)
+ else:
+ if (
+ attention_mask is not None
+ and attention_mask.dim() == 4
+ and attention_mask.dtype != torch.bool
+ ):
+ attention_mask = attention_mask.eq(0)
+ if SDPA_HAS_GQA:
+ A = scaled_dot_product_attention(
+ Qn,
+ Kn,
+ Vn,
+ attn_mask = attention_mask,
+ scale = self.scaling,
+ enable_gqa = True,
+ )
+ else:
+ A = scaled_dot_product_attention(
+ Qn,
+ Kn,
+ Vn,
+ attn_mask = attention_mask,
+ scale = self.scaling,
+ )
A = A.transpose(1, 2)
A = A.reshape(bsz, 1, attention_size)
A = fast_linear_forward(self.o_proj, A, out = self.temp_O)
@@ -442,6 +460,9 @@ def GraniteModel_fast_forward_inference(
hidden_states,
seq_len,
)
+ # Pre-convert to bool once for all layers (avoids per-layer .eq(0))
+ if attention_mask is not None and attention_mask.dtype != torch.bool:
+ attention_mask = attention_mask.eq(0)
else:
attention_mask = None
diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py
index 043d2363c1..93d93e26d6 100644
--- a/unsloth/models/llama.py
+++ b/unsloth/models/llama.py
@@ -25,6 +25,8 @@ from ._utils import move_to_device
from ._utils import (
_get_inference_mode_context_manager,
_prepare_model_for_qat,
+ is_bfloat16_supported,
+ get_quant_type,
)
from .loader_utils import _get_fp8_mode_and_check_settings
from ..utils.packing import (
@@ -35,6 +37,7 @@ from ..utils.attention_dispatch import (
AttentionConfig,
AttentionContext,
run_attention,
+ SDPA,
select_attention_backend,
)
from torch.nn.functional import scaled_dot_product_attention
@@ -212,11 +215,22 @@ def _fast_prepare_inputs_for_generation(
**kwargs,
):
past_key_values = kwargs.get("past_key_values", None)
+ original_attention_mask = attention_mask
# Handle inputs_embeds - only use on FIRST generation step (no cache)
# This fixes GitHub issue #3798: inputs_embeds was ignored
use_inputs_embeds = inputs_embeds is not None and past_key_values is None
+ if input_ids is not None and input_ids.numel() > 0:
+ bs, seq_length = input_ids.shape
+ device = input_ids.device
+ elif inputs_embeds is not None:
+ bs, seq_length, _ = inputs_embeds.shape
+ device = inputs_embeds.device
+ else:
+ bs, seq_length = 1, 0
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+
if past_key_values is not None:
# Check for uninitialized DynamicCache
if len(past_key_values) == 0:
@@ -233,16 +247,47 @@ def _fast_prepare_inputs_for_generation(
use_inputs_embeds = inputs_embeds is not None
else:
if input_ids is not None and input_ids.numel() > 0:
- bs, cache_length = input_ids.shape
+ bs = input_ids.shape[0]
input_ids = input_ids[:, [-1]]
device = input_ids.device
+ seq_length = 1
elif inputs_embeds is not None:
- bs, cache_length, _ = inputs_embeds.shape
+ bs, seq_length, _ = inputs_embeds.shape
device = inputs_embeds.device
else:
- bs, cache_length = 1, 0
+ bs, seq_length = 1, 0
device = "cuda" if torch.cuda.is_available() else "cpu"
+ if hasattr(past_key_values, "get_seq_length"):
+ past_len = int(past_key_values.get_seq_length())
+ else:
+ # legacy tuple cache: (layer, (K,V))
+ past_len = int(past_key_values[0][0].shape[-2])
+
+ max_cache_len = None
+ if hasattr(past_key_values, "get_max_cache_shape"):
+ m = past_key_values.get_max_cache_shape()
+ max_cache_len = int(m) if m is not None and m > 0 else None
+ elif hasattr(past_key_values, "get_max_length"):
+ m = past_key_values.get_max_length()
+ max_cache_len = int(m) if m is not None else None
+
+ # ensure cache_position
+ cache_position = kwargs.get("cache_position", None)
+ if cache_position is None:
+ kwargs["cache_position"] = torch.arange(
+ past_len,
+ past_len + seq_length,
+ device = device,
+ dtype = torch.long,
+ )
+ else:
+ if (
+ hasattr(cache_position, "device")
+ and cache_position.device != device
+ ):
+ kwargs["cache_position"] = cache_position.to(device)
+
# Get to the base model
base_model = self
if hasattr(base_model, "base_model_prefix"):
@@ -251,45 +296,49 @@ def _fast_prepare_inputs_for_generation(
if hasattr(
base_model, "_prepare_4d_causal_attention_mask_with_cache_position"
):
+ if not hasattr(base_model, "_unsloth_mask_needs_device"):
- def needs_device_kw(fn) -> bool:
- try:
- sig = inspect.signature(inspect.unwrap(fn))
- return "device" in sig.parameters
- except:
- # transformers <= 4.51.3 includes device arg but > 4.51.3 does not
- return transformers_version < Version("4.52.0")
+ def _check_needs_device(fn) -> bool:
+ try:
+ sig = inspect.signature(inspect.unwrap(fn))
+ return "device" in sig.parameters
+ except:
+ # transformers <= 4.51.3 includes device arg but > 4.51.3 does not
+ return transformers_version < Version("4.52.0")
- kwargs = {
- "sequence_length": 1,
- "target_length": cache_length,
+ base_model._unsloth_mask_needs_device = _check_needs_device(
+ base_model._prepare_4d_causal_attention_mask_with_cache_position
+ )
+
+ if max_cache_len is not None:
+ target_length = max_cache_len
+ elif (
+ original_attention_mask is not None
+ and original_attention_mask.dim() == 2
+ ):
+ target_length = original_attention_mask.shape[-1]
+ else:
+ target_length = past_len + seq_length
+
+ mask_kwargs = {
+ "sequence_length": seq_length,
+ "target_length": target_length,
"dtype": self.dtype,
- "cache_position": torch.arange(
- cache_length, cache_length + 1, device = device
- ),
+ "cache_position": kwargs["cache_position"],
"batch_size": bs,
"config": self.config,
"past_key_values": past_key_values,
}
- try:
- if needs_device_kw(
- base_model._prepare_4d_causal_attention_mask_with_cache_position
- ):
- kwargs["device"] = device
- except:
- print(
- f"Unsloth: Could not inspect signature of {base_model._prepare_4d_causal_attention_mask_with_cache_position}"
- )
+ if base_model._unsloth_mask_needs_device:
+ mask_kwargs["device"] = device
attention_mask = (
base_model._prepare_4d_causal_attention_mask_with_cache_position(
attention_mask,
- **kwargs,
+ **mask_kwargs,
)
)
else:
- if attention_mask is not None:
- attention_mask = attention_mask[:, [-1]]
if transformers_version <= Version("4.52.4"):
logger.warning_once(
f"{self.__class__.__name__} has no `_prepare_4d_causal_attention_mask_with_cache_position` method "
@@ -298,8 +347,17 @@ def _fast_prepare_inputs_for_generation(
"issue on GitHub."
)
- if "cache_position" in kwargs:
- kwargs["position_ids"] = kwargs["cache_position"]
+ if kwargs.get("position_ids", None) is None:
+ if original_attention_mask is not None and original_attention_mask.dim() == 2:
+ position_ids = original_attention_mask.long().cumsum(-1) - 1
+ position_ids.masked_fill_(original_attention_mask == 0, 1)
+ position_ids = position_ids[:, -seq_length:]
+ kwargs["position_ids"] = position_ids
+ elif kwargs.get("cache_position", None) is not None:
+ cp = kwargs["cache_position"]
+ if cp.dim() == 1:
+ cp = cp.unsqueeze(0).expand(bs, -1)
+ kwargs["position_ids"] = cp
result = {
"attention_mask": attention_mask,
@@ -329,6 +387,7 @@ def LlamaAttention_fast_forward_inference(
position_ids,
do_prefill = False,
attention_mask = None,
+ rotary_seq_len = None,
):
"""
https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L406
@@ -397,7 +456,7 @@ def LlamaAttention_fast_forward_inference(
# Mistral Nemo 12b has weird dimensions
if attention_size != hidden_size:
- self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = device)
+ self.temp_O = torch.empty((bsz, 1, hidden_size), dtype = dtype, device = device)
else:
self.temp_O = self.temp_QA[1][:, :, :hidden_size]
@@ -434,10 +493,19 @@ def LlamaAttention_fast_forward_inference(
# Need to do it prior 2 steps before hitting full on short KV cache
# or else error
- self.rotary_emb.extend_rope_embedding(Vn, seq_len + 2)
- cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index)
- cos = cos[position_ids].unsqueeze(1)
- sin = sin[position_ids].unsqueeze(1)
+ # ensure correct shape
+ if position_ids.dim() == 1:
+ position_ids = position_ids[:, None]
+ position_ids = position_ids.to(Qn.device)
+
+ if rotary_seq_len is None:
+ rotary_seq_len = max(kv_seq_len, int(position_ids.max().item()) + 1)
+ self.rotary_emb.extend_rope_embedding(Vn, rotary_seq_len + 1) # +1 slack
+ cos, sin = self.rotary_emb.get_cached(rotary_seq_len, Qn.device.index or 0)
+
+ cos = cos[position_ids].unsqueeze(1).to(device = Qn.device, dtype = Qn.dtype)
+ sin = sin[position_ids].unsqueeze(1).to(device = Qn.device, dtype = Qn.dtype)
+
h = self.half_head_dim
RH_Q = self.RH_Q
@@ -468,15 +536,17 @@ def LlamaAttention_fast_forward_inference(
sliding_window = getattr(self.config, "sliding_window", None)
if sliding_window is not None and kv_seq_len > sliding_window:
# From https://github.com/huggingface/transformers/blob/main/src/transformers/models/mistral/modeling_mistral.py#L193
- slicing_tokens = 1 - sliding_window
- Knn = Kn[:, :, slicing_tokens:, :] # .contiguous()
- Vnn = Vn[:, :, slicing_tokens:, :] # .contiguous()
+ start = kv_seq_len - sliding_window
+ Knn = Kn[:, :, start:, :] # .contiguous()
+ Vnn = Vn[:, :, start:, :] # .contiguous()
+ if attention_mask is not None:
+ attention_mask = attention_mask[..., start:]
else:
Knn, Vnn = Kn, Vn
# Grouped query attention
_, _, cached_len, _ = Knn.shape
- if bsz == 1 or not SDPA_HAS_GQA and n_groups != 1:
+ if bsz == 1 or ((not SDPA_HAS_GQA) and n_groups != 1):
Knn = Knn[:, :, None, :, :].expand(
bsz, n_kv_heads, n_groups, cached_len, head_dim
)
@@ -485,9 +555,6 @@ def LlamaAttention_fast_forward_inference(
)
Knn = Knn.reshape(bsz, n_heads, cached_len, head_dim)
Vnn = Vnn.reshape(bsz, n_heads, cached_len, head_dim)
- # else:
- # Knn, Vnn = Knn, Vnn
- # pass
# when qlen==vlen and attn_mask is None, we should use causal attention
Q_len = Qn.shape[-2]
@@ -503,12 +570,23 @@ def LlamaAttention_fast_forward_inference(
A = torch_matmul(
Qn, Knn.transpose(2, 3), out = self.attention[:, :, :, :cached_len]
)
- # if attention_mask is not None: A += attention_mask # Must add attention_mask for batched
A[:] = torch_nn_functional_softmax(
A, dim = -1, dtype = torch.float32
) # .to(A.dtype)
A = torch_matmul(A, Vnn, out = Qn)
+ # --- attention_mask fixup for SDPA if user passes 2D padding mask
else:
+ if attention_mask is not None and attention_mask.dim() == 2:
+ attention_mask = attention_mask[:, None, None, :].to(torch.bool)
+ # is it more appropriate to use _prepare_4d_causal_attention_mask_for_sdpa?
+ elif (
+ attention_mask is not None
+ and attention_mask.dim() == 4
+ and attention_mask.dtype != torch.bool
+ ):
+ # Decode is more stable with boolean keep masks than additive bf16 masks.
+ attention_mask = attention_mask.eq(0)
+
if SDPA_HAS_GQA:
A = scaled_dot_product_attention(
Qn,
@@ -662,6 +740,8 @@ def LlamaAttention_fast_forward(
rotary_emb = self.rotary_emb
rotary_emb.extend_rope_embedding(V, seq_len = kv_seq_len)
cos, sin = rotary_emb.get_cached(kv_seq_len, Q.device.index)
+ cos = cos.to(device = Q.device, dtype = Q.dtype)
+ sin = sin.to(device = Q.device, dtype = Q.dtype)
rope_position_ids = position_ids
if rope_position_ids is None and seq_info is not None:
@@ -681,7 +761,11 @@ def LlamaAttention_fast_forward(
# Attention module
use_varlen = seq_info is not None and past_key_value is None
- backend = select_attention_backend(use_varlen)
+ backend = (
+ SDPA if attention_mask is not None else select_attention_backend(use_varlen)
+ )
+
+ # should dropout be hardcoded to 0.0?
config = AttentionConfig(
backend = backend,
n_kv_heads = n_kv_heads,
@@ -1256,7 +1340,8 @@ def _LlamaModel_fast_forward_inference(
)
seq_len = past_key_values[0][0].shape[-2]
- if bsz != 1:
+ kv_seq_len = seq_len + 1
+ if attention_mask is not None:
attention_mask = _prepare_4d_causal_attention_mask_for_sdpa(
attention_mask,
(bsz, q_len),
@@ -1264,9 +1349,15 @@ def _LlamaModel_fast_forward_inference(
seq_len,
sliding_window = getattr(self.config, "sliding_window", None),
)
+ # Pre-convert to bool once for all layers (avoids per-layer .eq(0))
+ if attention_mask is not None and attention_mask.dtype != torch.bool:
+ attention_mask = attention_mask.eq(0)
else:
attention_mask = None
+ # Compute rotary_seq_len once to avoid per-layer GPU-CPU sync from .item()
+ rotary_seq_len = max(kv_seq_len, int(position_ids.max().item()) + 1)
+
next_decoder_cache = []
for idx, decoder_layer in enumerate(self.model.layers):
@@ -1289,6 +1380,7 @@ def _LlamaModel_fast_forward_inference(
position_ids = position_ids,
attention_mask = attention_mask,
do_prefill = not hasattr(decoder_layer.self_attn, "paged_attention"),
+ rotary_seq_len = rotary_seq_len,
)
X += residual
@@ -1527,7 +1619,7 @@ def CausalLM_fast_forward(fast_forward_inference):
logits = logit_softcapping * logits
else:
logits *= 1.0 / logit_softcapping
- torch.tanh(logits, out = logits)
+ logits.tanh_()
logits *= logit_softcapping
if not return_dict:
@@ -1640,6 +1732,17 @@ class LlamaRotaryEmbedding(torch.nn.Module):
self.multi_gpu_cos_cached = [None] * DEVICE_COUNT
self.multi_gpu_sin_cached = [None] * DEVICE_COUNT
+ # Normal Llama-3 RoPE
+ inv_freq = 1.0 / (
+ self.base
+ ** (
+ torch.arange(0, self.dim, 2, dtype = torch.int64, device = "cpu").float()
+ / self.dim
+ )
+ )
+ inv_freq = self._apply_inv_freq_scaling(inv_freq)
+ self.register_buffer("inv_freq", inv_freq, persistent = False)
+
# Build here to make `torch.jit.trace` work.
for device_idx in range(DEVICE_COUNT):
self._set_cos_sin_cache(
@@ -1656,22 +1759,24 @@ class LlamaRotaryEmbedding(torch.nn.Module):
1, device = get_current_device(), dtype = torch.get_default_dtype()
)
+ def _apply_inv_freq_scaling(self, inv_freq):
+ """Override to apply custom inv_freq scaling (e.g., extended RoPE)."""
+ return inv_freq
+
+ def _apply_time_scaling(self, t):
+ """Override to apply custom time scaling (e.g., linear scaling)."""
+ return t
+
def _set_cos_sin_cache(self, seq_len, device, dtype):
# Note: on the original Llama codebase, these tensors are created on the target device (and not on CPU) and
# in FP32. They are applied (multiplied) in FP32 as well.
self.current_rope_size = seq_len
- inv_freq = 1.0 / (
- self.base
- ** (
- torch.arange(0, self.dim, 2, dtype = torch.int64, device = "cpu").float()
- / self.dim
- )
- )
t = torch.arange(
- self.current_rope_size, device = "cpu", dtype = torch.int64
+ self.current_rope_size, device = self.inv_freq.device, dtype = torch.int64
).float()
+ t = self._apply_time_scaling(t)
- freqs = torch.outer(t, inv_freq)
+ freqs = torch.outer(t, self.inv_freq)
# Different from paper, but it uses a different permutation in order to obtain the same calculation
emb = torch.cat((freqs, freqs), dim = -1)
cos = emb.cos().to(dtype = dtype, device = device, non_blocking = True)
@@ -1733,33 +1838,14 @@ class LlamaLinearScalingRotaryEmbedding(LlamaRotaryEmbedding):
config = config,
)
- def _set_cos_sin_cache(self, seq_len, device, dtype):
- self.current_rope_size = seq_len
- inv_freq = 1.0 / (
- self.base
- ** (
- torch.arange(0, self.dim, 2, dtype = torch.int64, device = "cpu").float()
- / self.dim
- )
- )
- t = torch.arange(
- self.current_rope_size, device = "cpu", dtype = torch.int64
- ).float()
- t = t / self.scaling_factor
-
- freqs = torch.outer(t, inv_freq)
- # Different from paper, but it uses a different permutation in order to obtain the same calculation
- emb = torch.cat((freqs, freqs), dim = -1)
- cos = emb.cos().to(dtype = dtype, device = device, non_blocking = True)
- sin = emb.sin().to(dtype = dtype, device = device, non_blocking = True)
- self.multi_gpu_cos_cached[device.index] = cos
- self.multi_gpu_sin_cached[device.index] = sin
- return cos, sin
+ def _apply_time_scaling(self, t):
+ """Apply linear scaling to time indices."""
+ return t / self.scaling_factor
# See https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/rotary_embedding.py#L736
# For Llama 3.1
-class LlamaExtendedRotaryEmbedding(torch.nn.Module):
+class LlamaExtendedRotaryEmbedding(LlamaRotaryEmbedding):
def __init__(
self,
dim = None,
@@ -1768,101 +1854,16 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module):
device = None,
config = None, # [TODO] Hack to pass in config - need to remove later
):
- super().__init__()
- if config is not None:
- # [TODO] Hack to pass in config - need to remove later
- base = _get_rope_theta(config, default = base)
- partial_rotary_factor = (
- config.partial_rotary_factor
- if hasattr(config, "partial_rotary_factor")
- else 1.0
- )
- dim = int((config.hidden_size // config.num_attention_heads))
- device = DEVICE_TYPE_TORCH
- max_position_embeddings = config.max_position_embeddings
-
- self.dim = dim
- self.max_position_embeddings = max_position_embeddings
- self.base = base
- # Dynamic RoPE we first set it to a max of 4 * 8192 tokens then we iteratively grow this
- self.current_rope_size = min(4 * 8192, self.max_position_embeddings)
- self.multi_gpu_cos_cached = [None] * DEVICE_COUNT
- self.multi_gpu_sin_cached = [None] * DEVICE_COUNT
-
- # Normal Llama-3 RoPE
- inv_freq = 1.0 / (
- self.base
- ** (
- torch.arange(0, self.dim, 2, dtype = torch.int64, device = "cpu").float()
- / self.dim
- )
+ super().__init__(
+ dim = dim,
+ max_position_embeddings = max_position_embeddings,
+ base = base,
+ device = device,
+ config = config,
)
- inv_freq = self.apply_scaling(inv_freq)
- self.register_buffer("inv_freq", inv_freq, persistent = False)
-
- # Build here to make `torch.jit.trace` work.
- for device_idx in range(DEVICE_COUNT):
- self._set_cos_sin_cache(
- seq_len = self.current_rope_size,
- device = torch.device(device_idx),
- dtype = torch.get_default_dtype(),
- )
-
- # dummy so that patch_utils doesn't fail for now
- self.cos_cached = torch.empty(
- 1, device = get_current_device(), dtype = torch.get_default_dtype()
- )
- self.sin_cached = torch.empty(
- 1, device = get_current_device(), dtype = torch.get_default_dtype()
- )
-
- def _set_cos_sin_cache(self, seq_len, device, dtype):
- # Note: on the original Llama codebase, these tensors are created on the target device (and not on CPU) and
- # in FP32. They are applied (multiplied) in FP32 as well.
- self.current_rope_size = seq_len
-
- t = torch.arange(
- self.current_rope_size, device = self.inv_freq.device, dtype = torch.int64
- ).float()
-
- freqs = torch.outer(t, self.inv_freq)
- # Different from paper, but it uses a different permutation in order to obtain the same calculation
- emb = torch.cat((freqs, freqs), dim = -1)
- cos = emb.cos().to(dtype = dtype, device = device, non_blocking = True)
- sin = emb.sin().to(dtype = dtype, device = device, non_blocking = True)
- self.multi_gpu_cos_cached[device.index] = cos
- self.multi_gpu_sin_cached[device.index] = sin
- return cos, sin
-
- def forward(self, x, position_ids = None, seq_len = None):
- # x: [bs, num_attention_heads, seq_len, head_size]
- if seq_len is not None and seq_len > self.current_rope_size:
- self._set_cos_sin_cache(seq_len = seq_len, device = x.device, dtype = x.dtype)
- device_index = x.device.index
- return (
- self.multi_gpu_cos_cached[device_index][:seq_len],
- self.multi_gpu_sin_cached[device_index][:seq_len],
- )
-
- def get_cached(self, seq_len = None, device_index = None):
- if device_index is None:
- device_index = get_current_device()
- return self.multi_gpu_cos_cached[device_index], self.multi_gpu_sin_cached[
- device_index
- ]
-
- def extend_rope_embedding(self, x, seq_len):
- if seq_len <= self.current_rope_size:
- return
- # Iteratively grow by increments of 8192
- self.current_rope_size = ((seq_len // 8192) + ((seq_len % 8192) != 0)) * 8192
- for device_idx in range(DEVICE_COUNT):
- self._set_cos_sin_cache(
- self.current_rope_size, device = torch.device(device_idx), dtype = x.dtype
- )
# From https://github.com/meta-llama/llama-models/blob/main/models/llama3_1/api/model.py#L41
- def apply_scaling(self, freqs: torch.Tensor):
+ def _apply_inv_freq_scaling(self, freqs: torch.Tensor):
# Values obtained from grid search
scale_factor = 8
low_freq_factor = 1
@@ -2062,6 +2063,21 @@ def unsloth_fast_generate(
FastLlamaModel.for_inference(self)
+ # Unpack BatchEncoding passed as input_ids for backwards compatibility.
+ # Old notebooks do model.generate(input_ids=tokenizer(...)) where the tokenizer
+ # output is a BatchEncoding (dict-like). Transformers v5 generate() calls
+ # .shape on it directly and crashes. Unpack into separate kwargs so both
+ # v4 and v5 work transparently.
+ _maybe_encoding = kwargs.get("input_ids", None)
+ if (
+ _maybe_encoding is not None
+ and not isinstance(_maybe_encoding, torch.Tensor)
+ and hasattr(_maybe_encoding, "items")
+ ):
+ batch_data = kwargs.pop("input_ids")
+ for key, val in batch_data.items():
+ kwargs.setdefault(key, val)
+
dtype = _get_dtype(dtype_from_config(self.config))
if hasattr(self, "config") and hasattr(self.config, "max_position_embeddings"):
@@ -2071,9 +2087,6 @@ def unsloth_fast_generate(
and "max_new_tokens" in kwargs
):
_ids = kwargs["input_ids"]
- # Handle BatchEncoding from transformers 5.0+ (no .shape attribute)
- if hasattr(_ids, "input_ids"):
- _ids = _ids["input_ids"]
if hasattr(_ids, "shape") and (
_ids.shape[-1] + kwargs["max_new_tokens"]
> self.config.max_position_embeddings
@@ -2245,9 +2258,7 @@ class FastLlamaModel:
vllm_version = ""
elif DEVICE_TYPE == "hip":
gpu_stats = torch.cuda.get_device_properties(0)
- gpu_stats_name = (
- gpu_stats.name + ". " if gpu_stats.name != "" else "AMD GPU Device. "
- )
+ gpu_stats_name = resolve_hip_gpu_stats_name(gpu_stats)
gpu_version = torch.version.hip
gpu_stats_snippet = f"ROCm Toolkit: {gpu_version}."
try:
@@ -2669,16 +2680,39 @@ class FastLlamaModel:
model._old_generate = model.generate
unsloth_fast_generate.__doc__ = model._old_generate.__doc__
model.generate = types.MethodType(unsloth_fast_generate, model)
- # Set weight[padding_idx] = 0
- with torch.no_grad():
- for name, module in model.named_modules():
- if type(module) is torch.nn.Embedding:
- if (
- getattr(module, "weight", None) is not None
- and getattr(module, "padding_idx", None) is not None
- ):
- if module.padding_idx < module.weight.shape[0]:
- module.weight[module.padding_idx] = 0
+ # Set weight[padding_idx] = 0 for embeddings that are NOT tied with the
+ # lm_head. When weights are tied, zeroing the padding row also zeros
+ # the corresponding lm_head row, forcing logit = 0 for the pad token.
+ # This is higher than the (negative) logits for real tokens in models
+ # like Gemma, causing the decoder to emit and produce gibberish.
+ # Skip entirely if eos_token == pad_token to avoid zeroing EOS embedding.
+ eos_token_id = (
+ getattr(tokenizer, "eos_token_id", None) if tokenizer is not None else None
+ )
+ pad_token_id = (
+ getattr(tokenizer, "pad_token_id", None) if tokenizer is not None else None
+ )
+ if tokenizer is not None and eos_token_id != pad_token_id:
+ lm_head = getattr(model, "lm_head", None)
+ lm_head_weight = (
+ getattr(lm_head, "weight", None) if lm_head is not None else None
+ )
+ with torch.no_grad():
+ for name, module in model.named_modules():
+ if type(module) is torch.nn.Embedding:
+ if (
+ getattr(module, "weight", None) is not None
+ and getattr(module, "padding_idx", None) is not None
+ ):
+ if module.padding_idx < module.weight.shape[0]:
+ # Skip if tied to lm_head
+ if (
+ lm_head_weight is not None
+ and module.weight.data_ptr()
+ == lm_head_weight.data_ptr()
+ ):
+ continue
+ module.weight[module.padding_idx] = 0
return model, tokenizer
@staticmethod
@@ -2915,6 +2949,7 @@ class FastLlamaModel:
accepted_modules = frozenset(
(
+ "lm_head",
"q_proj",
"k_proj",
"v_proj",
@@ -2933,18 +2968,7 @@ class FastLlamaModel:
train_embed_tokens = False
final_modules = []
for module in target_modules:
- if module == "lm_head":
- # logger.warning_once(
- # "Unsloth: `lm_head` should be placed in `modules_to_save` and not `target_modules`. "\
- # "Luckily, we shall do it for you!"
- # )
- train_lm_head = True
- if modules_to_save is None:
- modules_to_save = ["lm_head"]
- else:
- modules_to_save.append("lm_head")
-
- elif module == "embed_tokens":
+ if module == "embed_tokens":
# logger.warning_once(
# "Unsloth: `embed_tokens` should be placed in `modules_to_save` and not `target_modules`. "\
# "Luckily, we shall do it for you!"
@@ -2969,23 +2993,32 @@ class FastLlamaModel:
# Check if we added new tokens!
if hasattr(model, "_need_to_train_embeddings"):
- if not train_lm_head or not train_embed_tokens:
+ # Check if embed_tokens/lm_head are already being trained
+ # (either as LoRA targets in final_modules or via modules_to_save)
+ _embed_already_trained = (
+ train_embed_tokens or "embed_tokens" in final_modules
+ )
+ _lm_head_already_trained = train_lm_head or "lm_head" in final_modules
+ if not _lm_head_already_trained or not _embed_already_trained:
print(
"Unsloth: You added new tokens but did not specify if you wanted to "
"train the lm_head and embed_tokens.\nWe must turn it on for you."
)
- train_lm_head = True
- train_embed_tokens = True
- if modules_to_save is None:
- modules_to_save = ["embed_tokens"]
- else:
- modules_to_save.append("embed_tokens")
+ # Only add to modules_to_save if not already a LoRA target
+ if not _embed_already_trained:
+ train_embed_tokens = True
+ if modules_to_save is None:
+ modules_to_save = ["embed_tokens"]
+ elif "embed_tokens" not in modules_to_save:
+ modules_to_save.append("embed_tokens")
- if modules_to_save is None:
- modules_to_save = ["lm_head"]
- else:
- modules_to_save.append("lm_head")
+ if not _lm_head_already_trained:
+ train_lm_head = True
+ if modules_to_save is None:
+ modules_to_save = ["lm_head"]
+ elif "lm_head" not in modules_to_save:
+ modules_to_save.append("lm_head")
# Check for Llama-3
# if hasattr(model._saved_temp_tokenizer, "_using_llama3_template"):
diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py
index 4054f1b7f5..e7b975ceea 100644
--- a/unsloth/models/loader.py
+++ b/unsloth/models/loader.py
@@ -78,6 +78,9 @@ SUPPORTS_QWEN3_MOE = transformers_version >= Version("4.50.3")
SUPPORTS_FALCON_H1 = transformers_version >= Version("4.53.0")
SUPPORTS_GEMMA3N = transformers_version >= Version("4.53.0")
SUPPORTS_GPTOSS = transformers_version >= Version("4.55.0")
+# Transformers v5 meta-device loading corrupts non-persistent buffers (inv_freq).
+# See _fix_rope_inv_freq() below for details.
+_NEEDS_ROPE_FIX = transformers_version >= Version("5.0.0")
if SUPPORTS_GEMMA:
from .gemma import FastGemmaModel
if SUPPORTS_GEMMA2:
@@ -103,6 +106,7 @@ FORCE_FLOAT32 = [
"gemma3text", # Gemma3TextModel (EmbeddingGemma, standalone text-only Gemma3)
"gemma3n",
"gpt_oss",
+ "qwen3_5", # Qwen3.5 RMSNorm uses (1+w) pattern like Gemma3, overflows float16
]
global DISABLE_COMPILE_MODEL_NAMES
@@ -121,6 +125,100 @@ DISABLE_SDPA_MODEL_NAMES = [
]
+def _fix_rope_inv_freq(model):
+ """Fix inv_freq corruption caused by transformers v5 meta-device loading.
+
+ Transformers v5 initializes models on the meta device, then
+ _move_missing_keys_from_meta_to_device() (modeling_utils.py) replaces ALL
+ non-persistent buffers with torch.empty_like() -- uninitialized memory.
+
+ Vanilla transformers restores inv_freq via _init_weights() which checks for
+ hasattr(module, "original_inv_freq"). Unsloth's LlamaRotaryEmbedding and
+ subclasses do not have this attribute, so inv_freq stays corrupted. This
+ produces wrong positional encodings and causes 5-11x higher training loss.
+
+ This function recomputes inv_freq from the stored base and dim, applies
+ any model-specific scaling, and rebuilds the cos/sin caches.
+
+ Only runs on transformers >= 5.0.0. No-op on v4.
+ """
+ if not _NEEDS_ROPE_FIX:
+ return model
+
+ for name, module in model.named_modules():
+ # Unsloth's LlamaRotaryEmbedding and subclasses (Extended, LinearScaling,
+ # Granite). Native v5 rotary classes (Gemma3, etc.) have original_inv_freq
+ # which v5's _init_weights() uses to restore inv_freq, so they are fine.
+ if (
+ hasattr(module, "inv_freq")
+ and hasattr(module, "base")
+ and hasattr(module, "dim")
+ and hasattr(module, "_apply_inv_freq_scaling")
+ and hasattr(module, "multi_gpu_cos_cached")
+ ):
+ inv_freq = 1.0 / (
+ module.base
+ ** (
+ torch.arange(
+ 0, module.dim, 2, dtype = torch.int64, device = "cpu"
+ ).float()
+ / module.dim
+ )
+ )
+ inv_freq = module._apply_inv_freq_scaling(inv_freq)
+ module.inv_freq = inv_freq
+ for device_idx in range(len(module.multi_gpu_cos_cached)):
+ if module.multi_gpu_cos_cached[device_idx] is not None:
+ module._set_cos_sin_cache(
+ seq_len = module.current_rope_size,
+ device = torch.device(device_idx),
+ dtype = torch.get_default_dtype(),
+ )
+
+ # LongRopeRotaryEmbedding (Phi-3.5 style with short_inv_freq + long_inv_freq)
+ elif (
+ hasattr(module, "short_inv_freq")
+ and hasattr(module, "long_inv_freq")
+ and hasattr(module, "base")
+ and hasattr(module, "dim")
+ ):
+ config = getattr(model, "config", None)
+ rope_scaling = getattr(config, "rope_scaling", None) if config else None
+ if rope_scaling is not None:
+ short_factor = rope_scaling.get("short_factor", None)
+ long_factor = rope_scaling.get("long_factor", None)
+ if short_factor is not None and long_factor is not None:
+ inv_freq_shape = (
+ torch.arange(
+ 0, module.dim, 2, dtype = torch.int64, device = "cpu"
+ ).float()
+ / module.dim
+ )
+ sf = torch.tensor(short_factor, device = "cpu", dtype = torch.float32)
+ lf = torch.tensor(long_factor, device = "cpu", dtype = torch.float32)
+ module.short_inv_freq = 1.0 / (sf * module.base**inv_freq_shape)
+ module.long_inv_freq = 1.0 / (lf * module.base**inv_freq_shape)
+
+ dtype = torch.bfloat16 if is_bfloat16_supported() else torch.float16
+ t = torch.arange(
+ module.original_max_position_embeddings,
+ device = module.short_inv_freq.device,
+ dtype = torch.int64,
+ ).float()
+ freqs = torch.outer(t, module.short_inv_freq)
+ emb = torch.cat((freqs, freqs), dim = -1)
+ for device_idx in range(len(module.multi_gpu_short_cos_cached)):
+ if module.multi_gpu_short_cos_cached[device_idx] is not None:
+ device_obj = torch.device(device_idx)
+ module.multi_gpu_short_cos_cached[device_idx] = (
+ emb.cos() * module.scaling_factor
+ ).to(dtype = dtype, device = device_obj, non_blocking = True)
+ module.multi_gpu_short_sin_cached[device_idx] = (
+ emb.sin() * module.scaling_factor
+ ).to(dtype = dtype, device = device_obj, non_blocking = True)
+ return model
+
+
class FastLanguageModel(FastLlamaModel):
@staticmethod
def from_pretrained(
@@ -261,12 +359,6 @@ class FastLanguageModel(FastLlamaModel):
fast_inference = False
break
- # [TODO] For now fast_inference only works with fast_inference ie vLLM
- if load_in_fp8 != False:
- if not fast_inference:
- raise NotImplementedError(
- "Unsloth: set `fast_inference = True` when doing `load_in_fp8`."
- )
# Check if 4bit is allowed specifically for AMD
if not ALLOW_BITSANDBYTES and not use_exact_model_name:
if load_in_4bit or load_in_8bit or model_name.lower().endswith("-bnb-4bit"):
@@ -280,7 +372,11 @@ class FastLanguageModel(FastLlamaModel):
fp8_mode = None
if not use_exact_model_name:
new_model_name = get_model_name(
- model_name, load_in_4bit = load_in_4bit, load_in_fp8 = load_in_fp8
+ model_name,
+ load_in_4bit = load_in_4bit,
+ load_in_fp8 = load_in_fp8,
+ token = token,
+ trust_remote_code = trust_remote_code,
)
if new_model_name is None and load_in_fp8 != False:
fp8_mode = _get_fp8_mode_and_check_settings(
@@ -433,7 +529,13 @@ class FastLanguageModel(FastLlamaModel):
# Check base model again for PEFT
model_name = peft_config.base_model_name_or_path
if not use_exact_model_name:
- model_name = get_model_name(model_name, load_in_4bit)
+ model_name = get_model_name(
+ model_name,
+ load_in_4bit = load_in_4bit,
+ load_in_fp8 = load_in_fp8,
+ token = token,
+ trust_remote_code = trust_remote_code,
+ )
# Check if pre-quantized models are allowed
# For eg AMD Instinct GPUs need blocksize = 128, but our pre-quants are blocksize = 64
if not ALLOW_PREQUANTIZED_MODELS and model_name.lower().endswith(
@@ -691,6 +793,7 @@ class FastLanguageModel(FastLlamaModel):
if patch_tiled_mlp_choice != "0" or unsloth_tiled_mlp:
patch_tiled_mlp(model, patch_options_str = patch_tiled_mlp_choice)
+ model = _fix_rope_inv_freq(model)
return model, tokenizer
@@ -876,13 +979,6 @@ class FastModel(FastBaseModel):
fast_inference = False
break
- # [TODO] For now fast_inference only works with fast_inference ie vLLM
- if load_in_fp8 != False:
- if not fast_inference:
- raise NotImplementedError(
- "Unsloth: set `fast_inference = True` when doing `load_in_fp8`."
- )
-
# Find FP8, BnB 4bit, other mapped names
old_model_name = model_name
fp8_mode = None
@@ -1046,6 +1142,14 @@ class FastModel(FastBaseModel):
# Set norms to float32 since anyways they get upcasted to float32
# common in both gemma-3 and gemma-3n
os.environ["UNSLOTH_HIGH_PRECISION_LAYERNORM"] = "1"
+ # ROCm/HIP: Gemma3 compiled forward produces NaN on RDNA GPUs
+ # (gfx1100, gfx1101, gfx1102, gfx1150, gfx1151, etc.).
+ # Disable torch.compile for model forward; loss compilation is fine.
+ # See https://github.com/unslothai/unsloth/issues/3385
+ from unsloth.kernels.utils import is_rdna
+
+ if is_rdna():
+ os.environ["UNSLOTH_COMPILE_DISABLE"] = "partial"
# Cohere
elif "cohere2" in model_types_all and transformers_version < Version(
"4.50.0.dev0"
@@ -1084,6 +1188,15 @@ class FastModel(FastBaseModel):
";"
"os.environ['TRITON_F32_DEFAULT'] = 'ieee'"
)
+ elif "nemotron_h" in model_types_all:
+ # NemotronH (hybrid Mamba-2 + Transformer) uses same Mamba kernels as Falcon-H1
+ # Mamba kernels need float32 Triton precision
+ os.environ["UNSLOTH_FORCE_CUSTOM_DTYPE"] = (
+ "float16;torch.float32;torch.float16;"
+ "if name.endswith(('q_proj', 'k_proj', 'v_proj', 'o_proj', 'gate_proj', 'up_proj', 'down_proj', 'head')): module.to(torch.float16)"
+ ";"
+ "os.environ['TRITON_F32_DEFAULT'] = 'ieee'"
+ )
elif "gpt_oss" in model_types_all:
os.environ["UNSLOTH_DISABLE_STATIC_GENERATION"] = "1"
if not load_in_4bit:
@@ -1281,7 +1394,21 @@ class FastModel(FastBaseModel):
is_vlm = any(x.endswith("ForConditionalGeneration") for x in architectures)
is_vlm = is_vlm or hasattr(model_config, "vision_config")
if auto_model is None:
- auto_model = AutoModelForVision2Seq if is_vlm else AutoModelForCausalLM
+ if is_vlm:
+ # Check if the model's auto_map supports the VLM auto class.
+ # Some VL models (e.g. Nemotron-VL) only register AutoModelForCausalLM
+ # in their auto_map, not AutoModelForImageTextToText/AutoModelForVision2Seq.
+ _auto_map = getattr(model_config, "auto_map", {}) or {}
+ _vlm_class_name = AutoModelForVision2Seq.__name__
+ if (
+ "AutoModelForCausalLM" in _auto_map
+ and _vlm_class_name not in _auto_map
+ ):
+ auto_model = AutoModelForCausalLM
+ else:
+ auto_model = AutoModelForVision2Seq
+ else:
+ auto_model = AutoModelForCausalLM
load_in_4bit_kwargs = load_in_4bit
load_in_8bit_kwargs = load_in_8bit
@@ -1398,6 +1525,7 @@ class FastModel(FastBaseModel):
if patch_tiled_mlp_choice != "0" or unsloth_tiled_mlp:
patch_tiled_mlp(model, patch_options_str = patch_tiled_mlp_choice)
+ model = _fix_rope_inv_freq(model)
return model, tokenizer
diff --git a/unsloth/models/loader_utils.py b/unsloth/models/loader_utils.py
index 01d221c725..cf5af983a6 100644
--- a/unsloth/models/loader_utils.py
+++ b/unsloth/models/loader_utils.py
@@ -31,7 +31,6 @@ from .mapper import (
from transformers import __version__ as transformers_version
from unsloth.models._utils import TorchAOConfig
from unsloth_zoo.utils import Version
-from unsloth_zoo.vllm_utils import _get_torchao_fp8_config
import gc
transformers_version = Version(transformers_version)
@@ -49,6 +48,13 @@ BAD_MAPPINGS = {
}
+def _get_torchao_fp8_config(fp8_mode):
+ # Import lazily so an optional, broken vLLM install does not break plain `import unsloth`.
+ from unsloth_zoo.vllm_utils import _get_torchao_fp8_config as _impl
+
+ return _impl(fp8_mode)
+
+
def _get_env_int(keys):
for key in keys:
value = os.environ.get(key)
@@ -192,18 +198,42 @@ def _get_new_mapper():
return {}, {}, {}
-def get_model_name(model_name, load_in_4bit = True, load_in_fp8 = False):
- assert load_in_fp8 in (True, False, "block")
- new_model_name = __get_model_name(
+def _resolve_with_mappers(
+ model_name,
+ load_in_4bit,
+ load_in_fp8,
+ int_to_float,
+ float_to_int,
+ map_to_unsloth_16bit,
+):
+ return __get_model_name(
model_name = model_name,
load_in_4bit = load_in_4bit,
- INT_TO_FLOAT_MAPPER = INT_TO_FLOAT_MAPPER,
- FLOAT_TO_INT_MAPPER = FLOAT_TO_INT_MAPPER,
- MAP_TO_UNSLOTH_16bit = MAP_TO_UNSLOTH_16bit,
+ INT_TO_FLOAT_MAPPER = int_to_float,
+ FLOAT_TO_INT_MAPPER = float_to_int,
+ MAP_TO_UNSLOTH_16bit = map_to_unsloth_16bit,
load_in_fp8 = load_in_fp8,
FLOAT_TO_FP8_BLOCK_MAPPER = FLOAT_TO_FP8_BLOCK_MAPPER,
FLOAT_TO_FP8_ROW_MAPPER = FLOAT_TO_FP8_ROW_MAPPER,
)
+
+
+def get_model_name(
+ model_name,
+ load_in_4bit = True,
+ load_in_fp8 = False,
+ token = None,
+ trust_remote_code = False,
+):
+ assert load_in_fp8 in (True, False, "block")
+ new_model_name = _resolve_with_mappers(
+ model_name = model_name,
+ load_in_4bit = load_in_4bit,
+ load_in_fp8 = load_in_fp8,
+ int_to_float = INT_TO_FLOAT_MAPPER,
+ float_to_int = FLOAT_TO_INT_MAPPER,
+ map_to_unsloth_16bit = MAP_TO_UNSLOTH_16bit,
+ )
# In the rare case, we convert bad model names to other names
# For eg too large dynamic quants or MoEs
if (
@@ -222,15 +252,13 @@ def get_model_name(model_name, load_in_4bit = True, load_in_fp8 = False):
NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER, NEW_MAP_TO_UNSLOTH_16bit = (
_get_new_mapper()
)
- upgraded_model_name = __get_model_name(
+ upgraded_model_name = _resolve_with_mappers(
model_name = model_name,
load_in_4bit = load_in_4bit,
- INT_TO_FLOAT_MAPPER = NEW_INT_TO_FLOAT_MAPPER,
- FLOAT_TO_INT_MAPPER = NEW_FLOAT_TO_INT_MAPPER,
- MAP_TO_UNSLOTH_16bit = NEW_MAP_TO_UNSLOTH_16bit,
load_in_fp8 = load_in_fp8,
- FLOAT_TO_FP8_BLOCK_MAPPER = FLOAT_TO_FP8_BLOCK_MAPPER,
- FLOAT_TO_FP8_ROW_MAPPER = FLOAT_TO_FP8_ROW_MAPPER,
+ int_to_float = NEW_INT_TO_FLOAT_MAPPER,
+ float_to_int = NEW_FLOAT_TO_INT_MAPPER,
+ map_to_unsloth_16bit = NEW_MAP_TO_UNSLOTH_16bit,
)
if upgraded_model_name is not None:
raise NotImplementedError(
@@ -239,10 +267,11 @@ def get_model_name(model_name, load_in_4bit = True, load_in_fp8 = False):
'pip install --upgrade --no-cache-dir "unsloth[colab-new] @ git+https://github.com/unslothai/unsloth.git"\n'
'pip install --upgrade --no-cache-dir "git+https://github.com/unslothai/unsloth-zoo.git"\n'
)
- if load_in_fp8 != False:
- # Handle on the fly TorchAO FP8 quantization
- return new_model_name
- return new_model_name if new_model_name is not None else model_name
+
+ if new_model_name is None:
+ new_model_name = model_name
+
+ return new_model_name
def _offline_quantize_to_fp8(model_name: str, fp8_mode: str) -> str:
@@ -344,10 +373,6 @@ def _get_fp8_mode_and_check_settings(
raise ValueError(
f"Unsloth: `load_in_fp8` can only be 'row' or 'block', got '{fp8_mode}'"
)
- if not fast_inference:
- raise ValueError(
- "Unsloth: `load_in_fp8` is only supported for `fast_inference` for now"
- )
if full_finetuning:
raise ValueError(
"Unsloth: `load_in_fp8` is not compatible with full finetuning"
diff --git a/unsloth/models/mapper.py b/unsloth/models/mapper.py
index e7296291f1..f0f430eb7e 100644
--- a/unsloth/models/mapper.py
+++ b/unsloth/models/mapper.py
@@ -965,11 +965,18 @@ __INT_TO_FLOAT_MAPPER = \
"mistralai/Magistral-Small-2506",
"unsloth/Magistral-Small-2506-bnb-4bit",
),
- "unsloth/Mistral-Small-3.2-24B-Instruct-2506-unsloth-bnb-4bit" : (
- "unsloth/Mistral-Small-3.2-24B-Instruct-2506",
- "mistralai/Mistral-Small-3.2-24B-Instruct-2506",
- "unsloth/Mistral-Small-3.2-24B-Instruct-2506-bnb-4bit",
- ),
+ "unsloth/Mistral-Small-3.2-24B-Instruct-2506-unsloth-bnb-4bit" : {
+ "8" : (
+ "mistralai/Mistral-Small-3.2-24B-Instruct-2506",
+ "unsloth/Mistral-Small-3.2-24B-Instruct-2506-FP8",
+ "unsloth/Mistral-Small-3.2-24B-Instruct-2506-FP8",
+ ),
+ "16" : (
+ "unsloth/Mistral-Small-3.2-24B-Instruct-2506",
+ "mistralai/Mistral-Small-3.2-24B-Instruct-2506",
+ "unsloth/Mistral-Small-3.2-24B-Instruct-2506-bnb-4bit",
+ ),
+ },
"unsloth/gemma-3n-E4B-it-unsloth-bnb-4bit" : (
"unsloth/gemma-3n-E4B-it",
"google/gemma-3n-E4B-it",
@@ -1056,11 +1063,18 @@ __INT_TO_FLOAT_MAPPER = \
"mistralai/Magistral-Small-2507",
"unsloth/Magistral-Small-2507-bnb-4bit",
),
- "unsloth/Magistral-Small-2509-unsloth-bnb-4bit" : (
- "unsloth/Magistral-Small-2509",
- "mistralai/Magistral-Small-2509",
- "unsloth/Magistral-Small-2509-bnb-4bit",
- ),
+ "unsloth/Magistral-Small-2509-unsloth-bnb-4bit" : {
+ "8" : (
+ "mistralai/Magistral-Small-2509",
+ "unsloth/Magistral-Small-2509-FP8-Dynamic",
+ "unsloth/Magistral-Small-2509-FP8-Dynamic",
+ ),
+ "16" : (
+ "unsloth/Magistral-Small-2509",
+ "mistralai/Magistral-Small-2509",
+ "unsloth/Magistral-Small-2509-bnb-4bit",
+ ),
+ },
"unsloth/Apertus-70B-Instruct-2509-unsloth-bnb-4bit" : (
"unsloth/Apertus-70B-Instruct-2509",
"swiss-ai/Apertus-70B-2509",
@@ -1256,6 +1270,76 @@ __INT_TO_FLOAT_MAPPER = \
"google/functiongemma-270m-it",
"unsloth/functiongemma-270m-it-unsloth-bnb-4bit",
),
+ # Ministral 3 models
+ "unsloth/Ministral-3-3B-Instruct-2512-unsloth-bnb-4bit" : {
+ "8" : (
+ "mistralai/Ministral-3-3B-Instruct-2512",
+ "unsloth/Ministral-3-3B-Instruct-2512-FP8",
+ "unsloth/Ministral-3-3B-Instruct-2512-FP8",
+ ),
+ "16" : (
+ "unsloth/Ministral-3-3B-Instruct-2512",
+ "mistralai/Ministral-3-3B-Instruct-2512",
+ "unsloth/Ministral-3-3B-Instruct-2512-bnb-4bit",
+ ),
+ },
+ "unsloth/Ministral-3-3B-Base-2512-unsloth-bnb-4bit" : (
+ "unsloth/Ministral-3-3B-Base-2512",
+ "mistralai/Ministral-3-3B-Base-2512",
+ "unsloth/Ministral-3-3B-Base-2512-bnb-4bit",
+ ),
+ "unsloth/Ministral-3-3B-Reasoning-2512-unsloth-bnb-4bit" : (
+ "unsloth/Ministral-3-3B-Reasoning-2512",
+ "mistralai/Ministral-3-3B-Reasoning-2512",
+ "unsloth/Ministral-3-3B-Reasoning-2512-bnb-4bit",
+ ),
+ "unsloth/Ministral-3-8B-Instruct-2512-unsloth-bnb-4bit" : {
+ "8" : (
+ "mistralai/Ministral-3-8B-Instruct-2512",
+ "unsloth/Ministral-3-8B-Instruct-2512-FP8",
+ "unsloth/Ministral-3-8B-Instruct-2512-FP8",
+ ),
+ "16" : (
+ "unsloth/Ministral-3-8B-Instruct-2512",
+ "mistralai/Ministral-3-8B-Instruct-2512",
+ "unsloth/Ministral-3-8B-Instruct-2512-bnb-4bit",
+ ),
+ },
+ "unsloth/Ministral-3-8B-Base-2512-unsloth-bnb-4bit" : (
+ "unsloth/Ministral-3-8B-Base-2512",
+ "mistralai/Ministral-3-8B-Base-2512",
+ "unsloth/Ministral-3-8B-Base-2512-bnb-4bit",
+ ),
+ "unsloth/Ministral-3-8B-Reasoning-2512-unsloth-bnb-4bit" : (
+ "unsloth/Ministral-3-8B-Reasoning-2512",
+ "mistralai/Ministral-3-8B-Reasoning-2512",
+ "unsloth/Ministral-3-8B-Reasoning-2512-bnb-4bit",
+ ),
+ "unsloth/Ministral-3-14B-Instruct-2512-unsloth-bnb-4bit" : {
+ "8" : (
+ "mistralai/Ministral-3-14B-Instruct-2512",
+ "unsloth/Ministral-3-14B-Instruct-2512-FP8",
+ "unsloth/Ministral-3-14B-Instruct-2512-FP8",
+ ),
+ "16" : (
+ "unsloth/Ministral-3-14B-Instruct-2512",
+ "mistralai/Ministral-3-14B-Instruct-2512",
+ "unsloth/Ministral-3-14B-Instruct-2512-bnb-4bit",
+ ),
+ },
+ "unsloth/Ministral-3-14B-Base-2512-unsloth-bnb-4bit" : (
+ "unsloth/Ministral-3-14B-Base-2512",
+ "mistralai/Ministral-3-14B-Base-2512",
+ "unsloth/Ministral-3-14B-Base-2512-bnb-4bit",
+ ),
+ "unsloth/Ministral-3-14B-Reasoning-2512-unsloth-bnb-4bit" : (
+ "unsloth/Ministral-3-14B-Reasoning-2512",
+ "mistralai/Ministral-3-14B-Reasoning-2512",
+ "unsloth/Ministral-3-14B-Reasoning-2512-bnb-4bit",
+ ),
+ "unsloth/Kimi-K2-Instruct-BF16" : (
+ "unsloth/Kimi-K2-Instruct",
+ ),
}
INT_TO_FLOAT_MAPPER = {}
@@ -1264,6 +1348,19 @@ MAP_TO_UNSLOTH_16bit = {}
FLOAT_TO_FP8_BLOCK_MAPPER = {}
FLOAT_TO_FP8_ROW_MAPPER = {}
+
+def _add_with_lower(mapper, key, value):
+ if key is None:
+ return
+ mapper[key] = value
+ mapper[key.lower()] = value
+
+
+def _add_lower_only(mapper, key, value):
+ if key is None:
+ return
+ mapper[key.lower()] = value
+
for key, values in __INT_TO_FLOAT_MAPPER.items():
block, row = None, None
if type(values) is dict:
@@ -1274,21 +1371,24 @@ for key, values in __INT_TO_FLOAT_MAPPER.items():
float8_values = values["8"]
assert len(float8_values) == 3
official, block, row = float8_values
- FLOAT_TO_FP8_BLOCK_MAPPER[key.lower()] = block
- FLOAT_TO_FP8_ROW_MAPPER[key.lower()] = row
- FLOAT_TO_FP8_BLOCK_MAPPER[official.lower() + "-dynamic"] = block
- FLOAT_TO_FP8_ROW_MAPPER[official.lower()] = row
- FLOAT_TO_FP8_ROW_MAPPER[official.lower() + "-dynamic"] = row
- FLOAT_TO_FP8_BLOCK_MAPPER[float16_values[0]] = block
- FLOAT_TO_FP8_BLOCK_MAPPER[float16_values[0].lower()] = block
- FLOAT_TO_FP8_ROW_MAPPER[float16_values[0]] = block
- FLOAT_TO_FP8_ROW_MAPPER[float16_values[0].lower()] = block
- for k in float8_values:
- FLOAT_TO_FP8_BLOCK_MAPPER[k.lower()] = block
- FLOAT_TO_FP8_ROW_MAPPER[k.lower()] = row
- for k in float16_values:
- FLOAT_TO_FP8_BLOCK_MAPPER[k.lower()] = block
- FLOAT_TO_FP8_ROW_MAPPER[k.lower()] = row
+ _add_lower_only(FLOAT_TO_FP8_BLOCK_MAPPER, key, block)
+ _add_lower_only(FLOAT_TO_FP8_ROW_MAPPER, key, row)
+ _add_lower_only(FLOAT_TO_FP8_BLOCK_MAPPER, official + "-dynamic", block)
+ _add_lower_only(FLOAT_TO_FP8_ROW_MAPPER, official, row)
+ _add_lower_only(FLOAT_TO_FP8_ROW_MAPPER, official + "-dynamic", row)
+ for k in float8_values + float16_values:
+ _add_lower_only(FLOAT_TO_FP8_BLOCK_MAPPER, k, block)
+ _add_lower_only(FLOAT_TO_FP8_ROW_MAPPER, k, row)
+
+ if float8_values[1] is not None and float8_values[1].startswith("unsloth"):
+ for value in float8_values:
+ if value is not None:
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, value, float8_values[1])
+
+ for value in float8_values:
+ if value is not None:
+ FLOAT_TO_INT_MAPPER[value] = key
+ FLOAT_TO_INT_MAPPER[value.lower()] = key.lower()
values = float16_values
INT_TO_FLOAT_MAPPER[key] = values[0]
@@ -1298,27 +1398,16 @@ for key, values in __INT_TO_FLOAT_MAPPER.items():
# Map to Unsloth version for 16bit versions
if len(values) == 2:
if values[0].startswith("unsloth"):
- MAP_TO_UNSLOTH_16bit[values[1]] = values[0]
- MAP_TO_UNSLOTH_16bit[values[1].lower()] = values[0]
- if block is not None:
- MAP_TO_UNSLOTH_16bit[block] = values[0]
- MAP_TO_UNSLOTH_16bit[block.lower()] = values[0]
- if row is not None:
- MAP_TO_UNSLOTH_16bit[row] = values[0]
- MAP_TO_UNSLOTH_16bit[row.lower()] = values[0]
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, values[1], values[0])
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, block, values[0])
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, row, values[0])
elif len(values) == 3:
# Dynamic Unsloth quantization
if values[0].startswith("unsloth"):
- MAP_TO_UNSLOTH_16bit[values[1]] = values[0]
- MAP_TO_UNSLOTH_16bit[values[1].lower()] = values[0]
- MAP_TO_UNSLOTH_16bit[values[2]] = values[0]
- MAP_TO_UNSLOTH_16bit[values[2].lower()] = values[0]
- if block is not None:
- MAP_TO_UNSLOTH_16bit[block] = values[0]
- MAP_TO_UNSLOTH_16bit[block.lower()] = values[0]
- if row is not None:
- MAP_TO_UNSLOTH_16bit[row] = values[0]
- MAP_TO_UNSLOTH_16bit[row.lower()] = values[0]
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, values[1], values[0])
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, values[2], values[0])
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, block, values[0])
+ _add_with_lower(MAP_TO_UNSLOTH_16bit, row, values[0])
pass
# Get lowercased
diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py
index 5e893d2b6f..83e9ab9486 100644
--- a/unsloth/models/mistral.py
+++ b/unsloth/models/mistral.py
@@ -25,6 +25,7 @@ from ..utils.attention_dispatch import (
AttentionConfig,
AttentionContext,
run_attention,
+ SDPA,
select_attention_backend,
)
from .llama import (
@@ -115,7 +116,9 @@ def MistralAttention_fast_forward(
use_varlen = (
seq_info is not None and past_key_value is None and window_size == (-1, -1)
)
- backend = select_attention_backend(use_varlen)
+ backend = (
+ SDPA if attention_mask is not None else select_attention_backend(use_varlen)
+ )
attention_config = AttentionConfig(
backend = backend,
n_kv_heads = n_kv_heads,
@@ -216,13 +219,18 @@ def MistralForCausalLM_fast_forward(
bsz, 1, q_len, q_len
)
else:
- # attention_mask should be [bsz, 1, q_len, q_len] or broadcastable
- # Add causal mask to existing attention mask
if attention_mask.dim() == 2:
- # [bsz, seq_len] -> [bsz, 1, 1, seq_len]
- attention_mask = attention_mask[:, None, None, :]
- attention_mask = attention_mask.expand(bsz, 1, q_len, q_len)
- attention_mask = attention_mask + causal_mask_values[None, None, :, :]
+ # Convert 0/1 padding mask to additive format: 1->0 (keep), 0->-inf (mask)
+ padding_mask = torch.where(
+ attention_mask[:, None, None, :].bool(),
+ 0.0,
+ -torch.inf,
+ )
+ attention_mask = causal_mask_values[None, None, :, :] + padding_mask
+ else:
+ attention_mask = (
+ attention_mask + causal_mask_values[None, None, :, :]
+ )
attention_mask = attention_mask.to(
dtype = _get_dtype(dtype_from_config(self.config))
diff --git a/unsloth/models/qwen3.py b/unsloth/models/qwen3.py
index ea06016d72..b93dddb186 100644
--- a/unsloth/models/qwen3.py
+++ b/unsloth/models/qwen3.py
@@ -21,6 +21,7 @@ from ..utils.attention_dispatch import (
AttentionConfig,
AttentionContext,
run_attention,
+ SDPA,
select_attention_backend,
)
from .llama import (
@@ -139,7 +140,9 @@ def Qwen3Attention_fast_forward(
# Attention module
use_varlen = seq_info is not None and past_key_value is None
- backend = select_attention_backend(use_varlen)
+ backend = (
+ SDPA if attention_mask is not None else select_attention_backend(use_varlen)
+ )
attention_config = AttentionConfig(
backend = backend,
n_kv_heads = n_kv_heads,
@@ -181,6 +184,7 @@ def Qwen3Attention_fast_forward_inference(
position_ids,
do_prefill = False,
attention_mask = None,
+ **kwargs,
):
"""
https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L406
@@ -249,7 +253,7 @@ def Qwen3Attention_fast_forward_inference(
# Mistral Nemo 12b has weird dimensions
if attention_size != hidden_size:
- self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = device)
+ self.temp_O = torch.empty((bsz, 1, hidden_size), dtype = dtype, device = device)
else:
self.temp_O = self.temp_QA[1][:, :, :hidden_size]
@@ -329,24 +333,42 @@ def Qwen3Attention_fast_forward_inference(
# Handle sliding windows
sliding_window = getattr(self.config, "sliding_window", None)
if sliding_window is not None and kv_seq_len > sliding_window:
- # From https://github.com/huggingface/transformers/blob/main/src/transformers/models/mistral/modeling_mistral.py#L193
- slicing_tokens = 1 - sliding_window
- Knn = Kn[:, :, slicing_tokens:, :] # .contiguous()
- Vnn = Vn[:, :, slicing_tokens:, :] # .contiguous()
+ start = kv_seq_len - sliding_window
+ Knn = Kn[:, :, start:, :] # .contiguous()
+ Vnn = Vn[:, :, start:, :] # .contiguous()
+ if attention_mask is not None:
+ attention_mask = attention_mask[..., start:]
else:
Knn, Vnn = Kn, Vn
# when qlen==vlen and attn_mask is None, we should use causal attention
Q_len = Qn.shape[-2]
K_len = Knn.shape[-2]
+ if attention_mask is not None and attention_mask.dim() == 2:
+ attention_mask = attention_mask[:, None, None, :].to(torch.bool)
+ elif (
+ attention_mask is not None
+ and attention_mask.dim() == 4
+ and attention_mask.dtype != torch.bool
+ ):
+ attention_mask = attention_mask.eq(0)
if attention_mask is None and Q_len == K_len:
is_causal = True
else:
is_causal = False
+ use_sdpa_gqa = SDPA_HAS_GQA
+ if (
+ use_sdpa_gqa
+ and isinstance(attention_mask, torch.Tensor)
+ and attention_mask.dim() >= 3
+ and attention_mask.shape[0] > 1
+ ):
+ # Avoid SDPA GQA drift for batched masked decode.
+ use_sdpa_gqa = False
# Grouped query attention
_, _, cached_len, _ = Knn.shape
- if bsz == 1 or not SDPA_HAS_GQA and n_groups != 1:
+ if bsz == 1 or ((not use_sdpa_gqa) and n_groups != 1):
Knn = Knn[:, :, None, :, :].expand(
bsz, n_kv_heads, n_groups, cached_len, head_dim
)
@@ -355,9 +377,6 @@ def Qwen3Attention_fast_forward_inference(
)
Knn = Knn.reshape(bsz, n_heads, cached_len, head_dim)
Vnn = Vnn.reshape(bsz, n_heads, cached_len, head_dim)
- # else:
- # Knn, Vnn = Knn, Vnn
- # pass
# Attention
if bsz == 1:
@@ -366,13 +385,12 @@ def Qwen3Attention_fast_forward_inference(
A = torch_matmul(
Qn, Knn.transpose(2, 3), out = self.attention[:, :, :, :cached_len]
)
- # if attention_mask is not None: A += attention_mask # Must add attention_mask for batched
A[:] = torch_nn_functional_softmax(
A, dim = -1, dtype = torch.float32
) # .to(A.dtype)
A = torch_matmul(A, Vnn, out = Qn)
else:
- if SDPA_HAS_GQA:
+ if use_sdpa_gqa:
A = scaled_dot_product_attention(
Qn,
Knn,
diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py
index 181e9479df..e4f34c908e 100755
--- a/unsloth/models/rl.py
+++ b/unsloth/models/rl.py
@@ -94,7 +94,40 @@ def vLLMSamplingParams(**kwargs):
def PatchRL(FastLanguageModel):
- from trl.models.utils import unwrap_model_for_generation
+ try:
+ from trl.models.utils import unwrap_model_for_generation
+ except ImportError:
+ try:
+ from trl.models import unwrap_model_for_generation
+ except ImportError:
+ # Local fallback -- TRL removed or moved this symbol
+ from contextlib import contextmanager as _cm
+
+ @_cm
+ def unwrap_model_for_generation(
+ model, accelerator, gather_deepspeed3_params = True
+ ):
+ unwrapped_model = accelerator.unwrap_model(model)
+ is_gc = getattr(unwrapped_model, "is_gradient_checkpointing", False)
+ if is_gc:
+ unwrapped_model.gradient_checkpointing_disable()
+ if (
+ getattr(accelerator, "state", None) is not None
+ and getattr(accelerator.state, "deepspeed_plugin", None) is not None
+ and accelerator.state.deepspeed_plugin.zero_stage == 3
+ ):
+ if not gather_deepspeed3_params:
+ yield accelerator.unwrap_model(model)
+ else:
+ import deepspeed
+
+ with deepspeed.zero.GatheredParameters(model.parameters()):
+ yield accelerator.unwrap_model(model)
+ else:
+ yield unwrapped_model
+ if is_gc:
+ unwrapped_model.gradient_checkpointing_enable()
+
from contextlib import contextmanager
@contextmanager
@@ -253,9 +286,12 @@ create_completion_attention_mask = RL_REPLACEMENTS["create_completion_attention_
left_pack_padding = RL_REPLACEMENTS["left_pack_padding"]
align_logprobs_with_mask = RL_REPLACEMENTS["align_logprobs_with_mask"]
autotune_batch_and_chunks = RL_REPLACEMENTS["grpo_autotune_batch_and_chunks"]
+sanitize_logprob = RL_REPLACEMENTS["sanitize_logprob"]
RLTrainer_replacement = '''
import os
+import math
+import logging
from typing import *
from dataclasses import dataclass, field
from packaging.version import Version
@@ -324,6 +360,7 @@ torch_compile_options = {{
{left_pack_padding_code}
{align_logprobs_with_mask_code}
{autotune_batch_and_chunks_code}
+{sanitize_logprob_code}
{RL_pre}
@@ -773,6 +810,18 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
" os.environ['UNSLOTH_RETURN_LOGITS'] = '1'\n"
)
extra_args += logits_check
+ warnings_issued_check = (
+ "if model is not None:\n"
+ " _warnings_issued = getattr(model, 'warnings_issued', None)\n"
+ " if _warnings_issued is None:\n"
+ " model.warnings_issued = {}\n"
+ " elif not isinstance(_warnings_issued, dict):\n"
+ " try:\n"
+ " model.warnings_issued = dict(_warnings_issued)\n"
+ " except Exception:\n"
+ " model.warnings_issued = {}\n"
+ )
+ extra_args += warnings_issued_check
# Check max_seq_length
if "model" in call_args:
@@ -988,6 +1037,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
"include_num_input_tokens_seen": False,
"auto_find_batch_size": False, # Auto /2 batch size - too many people complained so removing
"dataloader_pin_memory": True,
+ "padding_free": None, # None = user didn't set it, allows auto-enable detection
# Might fail so disable for now
# "dataloader_persistent_workers" : True, # Keeps dataloader in RAM
# "dataloader_prefetch_factor" : 2,
@@ -1228,6 +1278,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
left_pack_padding_code = inspect.getsource(left_pack_padding)
align_logprobs_with_mask_code = inspect.getsource(align_logprobs_with_mask)
autotune_batch_and_chunks_code = inspect.getsource(autotune_batch_and_chunks)
+ sanitize_logprob_code = inspect.getsource(sanitize_logprob)
# Get final source code
RLTrainer_source = RLTrainer_replacement.format(
RLTrainer_name = RLTrainer_name,
@@ -1256,6 +1307,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
autotune_batch_and_chunks_code = autotune_batch_and_chunks_code,
left_pack_padding_code = left_pack_padding_code,
align_logprobs_with_mask_code = align_logprobs_with_mask_code,
+ sanitize_logprob_code = sanitize_logprob_code,
)
if RLTrainer_name == "GRPOTrainer":
diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py
index 27f00f10f3..b3a55440f9 100755
--- a/unsloth/models/rl_replacements.py
+++ b/unsloth/models/rl_replacements.py
@@ -355,8 +355,9 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
re.DOTALL | re.MULTILINE,
)
+ # sanitize_logprob is injected as a module-level function via RLTrainer_replacement
+ # template in rl.py (from RL_REPLACEMENTS), so just reference it directly here.
replacement_text = (
- r"\1from trl.scripts.vllm_serve import sanitize_logprob\n"
r"\1all_logprobs = [\n"
r"\1 [sanitize_logprob(next(iter(logprob.values()))) for logprob in output.logprobs]\n"
r"\1 for outputs in all_outputs\n"
diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py
index c294dbdb0b..a8adba99e7 100644
--- a/unsloth/models/vision.py
+++ b/unsloth/models/vision.py
@@ -126,11 +126,6 @@ _compile_config = CompileConfig(
)
_compile_config.disable = True # Must set manually
-from unsloth_zoo.vllm_utils import (
- convert_lora_modules,
- return_lora_modules,
-)
-
try:
torch_compiler_set_stance = torch.compiler.set_stance
except:
@@ -497,9 +492,7 @@ class FastBaseModel:
vllm_version = ""
elif DEVICE_TYPE == "hip":
gpu_stats = torch.cuda.get_device_properties(0)
- gpu_stats_name = (
- gpu_stats.name + ". " if gpu_stats.name != "" else "AMD GPU Device. "
- )
+ gpu_stats_name = resolve_hip_gpu_stats_name(gpu_stats)
gpu_version = torch.version.hip
gpu_stats_snippet = f"ROCm Toolkit: {gpu_version}."
try:
@@ -618,7 +611,22 @@ class FastBaseModel:
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"
+ # Handle FP8 models: get_model_name has already redirected this to BF16 sibling if the model ships with
+ # FP8 weights. We just need to update it here for sanity.
+ auto_config.model_name = model_name
+ # Re-resolve model_class after potential config change
+ try:
+ model_class = auto_model._model_mapping[auto_config.__class__]
+ except Exception:
+ model_class = None
+
+ model_type = str(getattr(auto_config, "model_type", "")).lower()
+ if model_type.startswith("gemma3n"):
+ # Gemma3N variants initialize timm-based vision towers which do
+ # not support flex_attention, so default to eager unless overridden.
+ default_attn_impl = "eager"
+ else:
+ default_attn_impl = "flex_attention" if flex_attn_impl else "sdpa"
if not ("attn_implementation" in kwargs):
kwargs["attn_implementation"] = default_attn_impl
if not supports_sdpa and kwargs.get("attn_implementation") == "sdpa":
@@ -642,18 +650,25 @@ class FastBaseModel:
raise RuntimeError(
"Unsloth: Can only load in 4bit or 8bit or 16bit, not a combination!"
)
+ _skip_modules = SKIP_QUANTIZATION_MODULES.copy()
+ # Nemotron-H uses 'mixer' (not 'mamba') for Mamba layers.
+ # Mamba fused kernels pass out_proj.weight directly to F.linear,
+ # which fails with quantized Params4bit. Skip out_proj from quantization.
+ if any(mt == "nemotron_h" for mt in (model_types or [])):
+ _skip_modules.append("out_proj")
+
if load_in_4bit:
bnb_config = BitsAndBytesConfig(
load_in_4bit = True,
bnb_4bit_use_double_quant = True,
bnb_4bit_quant_type = "nf4",
bnb_4bit_compute_dtype = bnb_compute_dtype,
- llm_int8_skip_modules = SKIP_QUANTIZATION_MODULES.copy(),
+ llm_int8_skip_modules = _skip_modules,
)
elif load_in_8bit:
bnb_config = BitsAndBytesConfig(
load_in_8bit = True,
- llm_int8_skip_modules = SKIP_QUANTIZATION_MODULES.copy(),
+ llm_int8_skip_modules = _skip_modules,
)
elif load_in_16bit:
bnb_config = None
@@ -766,6 +781,7 @@ class FastBaseModel:
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()
@@ -774,6 +790,7 @@ class FastBaseModel:
load_in_fp8 = kwargs.pop("load_in_fp8", None)
model = auto_model.from_pretrained(
model_name,
+ config = model_config,
device_map = device_map,
# torch_dtype = torch_dtype, # Transformers removed torch_dtype
# quantization_config = bnb_config,
@@ -923,6 +940,32 @@ class FastBaseModel:
# Counteract saved tokenizers
tokenizer_name = model_name if tokenizer_name is None else tokenizer_name
+
+ # Fix _Unsloth_Patched_ prefix in local config files from old saves (issue #4085)
+ if os.path.isdir(tokenizer_name):
+ import json as _json
+
+ for _cfg_name in (
+ "processor_config.json",
+ "preprocessor_config.json",
+ "tokenizer_config.json",
+ ):
+ _cfg_path = os.path.join(tokenizer_name, _cfg_name)
+ if os.path.exists(_cfg_path):
+ try:
+ with open(_cfg_path, "r", encoding = "utf-8") as _f:
+ _cfg = _json.load(_f)
+ if _cfg.get("processor_class", "").startswith(
+ "_Unsloth_Patched_"
+ ):
+ _cfg["processor_class"] = _cfg["processor_class"][
+ len("_Unsloth_Patched_") :
+ ]
+ with open(_cfg_path, "w", encoding = "utf-8") as _f:
+ _json.dump(_cfg, _f, indent = 2, ensure_ascii = False)
+ except Exception:
+ pass
+
if (whisper_language and whisper_task) or auto_model.__name__.endswith(
"ForConditionalGeneration"
):
@@ -954,14 +997,23 @@ class FastBaseModel:
)
# If processor loading failed (e.g., tokenizer class not found),
+ # or if AutoProcessor silently degraded to a text-only tokenizer
+ # instead of returning a full VLM processor (issue #4085),
# try constructing the processor manually from separate components.
- if tokenizer is None and is_vlm:
- tokenizer = _construct_vlm_processor_fallback(
+ _processor_is_degraded = (
+ is_vlm
+ and tokenizer is not None
+ and not hasattr(tokenizer, "image_processor")
+ )
+ if (tokenizer is None or _processor_is_degraded) and is_vlm:
+ _fallback = _construct_vlm_processor_fallback(
tokenizer_name,
model_type_arch,
token,
trust_remote_code,
)
+ if _fallback is not None:
+ tokenizer = _fallback
if tokenizer is None:
import sys
@@ -992,6 +1044,7 @@ class FastBaseModel:
do_forced_float32 = do_forced_float32,
correct_dtype = correct_dtype,
)
+
try:
model, tokenizer = patch_tokenizer(model, tokenizer)
except Exception as _patch_err:
@@ -1354,9 +1407,15 @@ class FastBaseModel:
m.for_training = functools.partial(FastBaseModel.for_training, m)
m.for_inference = functools.partial(FastBaseModel.for_inference, m)
m = m.model
- # Set weight[padding_idx] = 0
+ # Set weight[padding_idx] = 0 for embeddings that are NOT tied with the
+ # lm_head. When weights are tied, zeroing the padding row also zeros
+ # the corresponding lm_head row, forcing logit = 0 for the pad token.
# Only do this if tokenizer is defined since eos_token == pad_token sometimes!
pad_token_id = getattr(tokenizer, "pad_token_id", None)
+ lm_head = getattr(model, "lm_head", None)
+ lm_head_weight = (
+ getattr(lm_head, "weight", None) if lm_head is not None else None
+ )
if (
tokenizer is not None
and getattr(tokenizer, "eos_token_id", None) != pad_token_id
@@ -1372,6 +1431,13 @@ class FastBaseModel:
module.padding_idx == pad_token_id
and module.padding_idx < module.weight.shape[0]
):
+ # Skip if tied to lm_head
+ if (
+ lm_head_weight is not None
+ and module.weight.data_ptr()
+ == lm_head_weight.data_ptr()
+ ):
+ continue
module.weight[module.padding_idx] = 0
return model
diff --git a/unsloth/save.py b/unsloth/save.py
index fc3b7b8771..6e38d1e952 100644
--- a/unsloth/save.py
+++ b/unsloth/save.py
@@ -23,6 +23,16 @@ from unsloth_zoo.llama_cpp import (
check_llama_cpp,
_download_convert_hf_to_gguf,
)
+
+# H4: Defensive imports -- these were added in unsloth-zoo PR #526
+# and may not exist on older versions
+try:
+ from unsloth_zoo.llama_cpp import LLAMA_CPP_DEFAULT_DIR, IS_WINDOWS
+except ImportError:
+ import sys
+
+ IS_WINDOWS = sys.platform == "win32"
+ LLAMA_CPP_DEFAULT_DIR = "llama.cpp"
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
from peft.tuners.lora import Linear as Peft_Linear
@@ -257,6 +267,7 @@ def unsloth_save_model(
# Our functions
temporary_location: str = "_unsloth_temporary_saved_buffers",
maximum_memory_usage: float = 0.9,
+ datasets: Optional[List[str]] = None,
):
if token is None:
token = get_token()
@@ -289,6 +300,7 @@ def unsloth_save_model(
"save_method",
"temporary_location",
"maximum_memory_usage",
+ "datasets",
):
del save_pretrained_settings[deletion]
@@ -366,6 +378,7 @@ def unsloth_save_model(
file_location = None,
old_username = None,
private = private,
+ datasets = datasets,
)
getattr(model, "original_push_to_hub", model.push_to_hub)(
@@ -475,6 +488,7 @@ def unsloth_save_model(
file_location = None,
old_username = None,
private = private,
+ datasets = datasets,
)
if tokenizer is not None:
@@ -737,6 +751,7 @@ def unsloth_save_model(
file_location = None,
old_username = username,
private = private,
+ datasets = datasets,
)
# First check if we're pushing to an organization!
@@ -1311,18 +1326,26 @@ def save_to_gguf(
"`model.{save_pretrained/push_to_hub}_gguf will use too much disk space.\n"
"You can try saving it to the `/tmp` directory for larger disk space.\n"
"I suggest you to save the 16bit model first, then use manual llama.cpp conversion.\n"
- "Error: {e}"
+ f"Error: {e}"
)
else:
+ if IS_WINDOWS:
+ build_instructions = (
+ f'cd "{LLAMA_CPP_DEFAULT_DIR}"\n'
+ f"cmake -S . -B build -DBUILD_SHARED_LIBS=OFF\n"
+ f"cmake --build build --config Release"
+ )
+ else:
+ build_instructions = f'cd "{LLAMA_CPP_DEFAULT_DIR}" && make clean && make all -j'
+
raise RuntimeError(
f"Unsloth: Quantization failed for {output_location}\n"
"You might have to compile llama.cpp yourself, then run this again.\n"
"You do not need to close this Python program. Run the following commands in a new terminal:\n"
- "You must run this in the same folder as you're saving your model.\n"
- "git clone --recursive https://github.com/ggerganov/llama.cpp\n"
- "cd llama.cpp && make clean && make all -j\n"
+ f'git clone --recursive https://github.com/ggerganov/llama.cpp "{LLAMA_CPP_DEFAULT_DIR}"\n'
+ f"{build_instructions}\n"
"Once that's done, redo the quantization.\n"
- "Error: {e}"
+ f"Error: {e}"
)
print("Unsloth: Model files cleanup...")
if quants_created:
@@ -1362,6 +1385,7 @@ def unsloth_save_pretrained_merged(
tags: List[str] = None,
temporary_location: str = "_unsloth_temporary_saved_buffers",
maximum_memory_usage: float = 0.75,
+ datasets: Optional[List[str]] = None,
):
"""
Same as .save_pretrained(...) except 4bit weights are auto
@@ -1403,6 +1427,7 @@ def unsloth_push_to_hub_merged(
tags: Optional[List[str]] = None,
temporary_location: str = "_unsloth_temporary_saved_buffers",
maximum_memory_usage: float = 0.75,
+ datasets: Optional[List[str]] = None,
):
"""
Same as .push_to_hub(...) except 4bit weights are auto
@@ -1480,10 +1505,11 @@ def create_huggingface_repo(
save_directory,
token = None,
private = False,
+ datasets = None,
):
if token is None:
token = get_token()
- save_directory, username = _determine_username(save_directory, "", token)
+ save_directory, username = _determine_username(save_directory, None, token)
from huggingface_hub import create_repo
@@ -1507,9 +1533,22 @@ def create_huggingface_repo(
extra = "unsloth",
)
card = ModelCard(content)
+ if datasets:
+ card.data.datasets = datasets
card.push_to_hub(save_directory, token = token)
except:
- pass
+ # Repo already exists — update datasets metadata separately
+ if datasets:
+ try:
+ from huggingface_hub import metadata_update
+
+ metadata_update(
+ save_directory, {"datasets": datasets}, overwrite = True, token = token
+ )
+ except Exception as e:
+ logger.warning_once(
+ f"Unsloth: Could not update datasets metadata for {save_directory}: {e}"
+ )
hf_api = HfApi(token = token)
return save_directory, hf_api
@@ -1524,6 +1563,7 @@ def upload_to_huggingface(
old_username = None,
private = None,
create_config = True,
+ datasets = None,
):
save_directory, username = _determine_username(save_directory, old_username, token)
@@ -1549,9 +1589,22 @@ def upload_to_huggingface(
extra = extra,
)
card = ModelCard(content)
+ if datasets:
+ card.data.datasets = datasets
card.push_to_hub(save_directory, token = token)
except:
- pass
+ # Repo already exists — update datasets metadata separately
+ if datasets:
+ try:
+ from huggingface_hub import metadata_update
+
+ metadata_update(
+ save_directory, {"datasets": datasets}, overwrite = True, token = token
+ )
+ except Exception as e:
+ logger.warning_once(
+ f"Unsloth: Could not update datasets metadata for {save_directory}: {e}"
+ )
if file_location is not None:
# Now upload file
@@ -1890,6 +1943,20 @@ def unsloth_save_pretrained_gguf(
arguments["push_to_hub"] = False # We handle upload ourselves
# GPT-OSS needs mxfp4 save method
if is_gpt_oss:
+ if quantization_method is not None:
+ _qm = (
+ quantization_method
+ if isinstance(quantization_method, (list, tuple))
+ else [quantization_method]
+ )
+ _ignored = [q for q in _qm if str(q).lower() != "mxfp4"]
+ if _ignored:
+ logger.warning_once(
+ f"Unsloth: GPT-OSS does not support GGUF quantization "
+ f"(requested: {', '.join(str(q) for q in _ignored)}). "
+ f"Overriding to MXFP4 format. "
+ f"Pass quantization_method=None to suppress this warning."
+ )
arguments["save_method"] = "mxfp4"
else:
arguments["save_method"] = "merged_16bit"
@@ -2035,16 +2102,22 @@ def unsloth_save_pretrained_gguf(
"Unsloth: ##### We removed it in GGUF's chat template for you."
)
+ _exe = ".exe" if IS_WINDOWS else ""
+ if IS_WINDOWS:
+ _bin_dir = os.path.join(LLAMA_CPP_DEFAULT_DIR, "build", "bin", "Release")
+ else:
+ _bin_dir = LLAMA_CPP_DEFAULT_DIR
+
if is_vlm_update:
print("\n")
print(
- f"Unsloth: example usage for Multimodal LLMs: llama.cpp/llama-mtmd-cli -m {all_file_locations[0]} --mmproj {all_file_locations[-1]}"
+ f"Unsloth: example usage for Multimodal LLMs: {os.path.join(_bin_dir, 'llama-mtmd-cli' + _exe)} -m {all_file_locations[0]} --mmproj {all_file_locations[-1]}"
)
print("Unsloth: load image inside llama.cpp runner: /image test_image.jpg")
print("Unsloth: Prompt model to describe the image")
else:
print(
- f'Unsloth: example usage for text only LLMs: llama.cpp/llama-cli --model {all_file_locations[0]} -p "why is the sky blue?"'
+ f'Unsloth: example usage for text only LLMs: {os.path.join(_bin_dir, "llama-cli" + _exe)} --model {all_file_locations[0]} -p "why is the sky blue?"'
)
if ollama_success:
@@ -2083,6 +2156,7 @@ def unsloth_push_to_hub_gguf(
tags: Optional[List[str]] = None,
temporary_location: str = "_unsloth_temporary_saved_buffers",
maximum_memory_usage: float = 0.85,
+ datasets: Optional[List[str]] = None,
):
"""
Same as .push_to_hub(...) except 4bit weights are auto
@@ -2257,8 +2331,8 @@ tags:
This model was finetuned and converted to GGUF format using [Unsloth](https://github.com/unslothai/unsloth).
**Example usage**:
-- For text only LLMs: `./llama.cpp/llama-cli -hf {repo_id} --jinja`
-- For multimodal models: `./llama.cpp/llama-mtmd-cli -hf {repo_id} --jinja`
+- For text only LLMs: `llama-cli -hf {repo_id} --jinja`
+- For multimodal models: `llama-mtmd-cli -hf {repo_id} --jinja`
## Available Model files:
"""
@@ -2338,6 +2412,18 @@ This model was finetuned and converted to GGUF format using [Unsloth](https://gi
except:
pass
+ if datasets:
+ try:
+ from huggingface_hub import metadata_update
+
+ metadata_update(
+ full_repo_id, {"datasets": datasets}, overwrite = True, token = token
+ )
+ except Exception as e:
+ logger.warning_once(
+ f"Unsloth: Could not update datasets metadata for {full_repo_id}: {e}"
+ )
+
except Exception as e:
raise RuntimeError(f"Failed to upload to Hugging Face Hub: {e}")
@@ -2645,6 +2731,7 @@ def unsloth_generic_save(
# Our functions
temporary_location: str = "_unsloth_temporary_saved_buffers",
maximum_memory_usage: float = 0.9,
+ datasets: Optional[List[str]] = None,
):
if token is None and push_to_hub:
token = get_token()
@@ -2672,6 +2759,20 @@ def unsloth_generic_save(
low_disk_space_usage = True,
use_temp_file = False,
)
+
+ if push_to_hub and datasets:
+ try:
+ from huggingface_hub import metadata_update
+
+ save_dir, _ = _determine_username(save_directory, None, token)
+ metadata_update(
+ save_dir, {"datasets": datasets}, overwrite = True, token = token
+ )
+ except Exception as e:
+ logger.warning_once(
+ f"Unsloth: Could not update datasets metadata for {save_directory}: {e}"
+ )
+
return
@@ -2692,6 +2793,7 @@ def unsloth_generic_save_pretrained_merged(
tags: List[str] = None,
temporary_location: str = "_unsloth_temporary_saved_buffers",
maximum_memory_usage: float = 0.75,
+ datasets: Optional[List[str]] = None,
):
"""
Same as .push_to_hub(...) except 4bit weights are auto
@@ -2733,6 +2835,7 @@ def unsloth_generic_push_to_hub_merged(
tags: Optional[List[str]] = None,
temporary_location: str = "_unsloth_temporary_saved_buffers",
maximum_memory_usage: float = 0.75,
+ datasets: Optional[List[str]] = None,
):
"""
Same as .push_to_hub(...) except 4bit weights are auto
diff --git a/unsloth/tokenizer_utils.py b/unsloth/tokenizer_utils.py
index 1c107bd84f..c445879df7 100644
--- a/unsloth/tokenizer_utils.py
+++ b/unsloth/tokenizer_utils.py
@@ -1007,7 +1007,14 @@ def patch_sft_trainer_tokenizer():
function = function.replace(replacer, check_text + replacer)
x = [x for x in all_imports if x in function]
- exec(f"from trl.trainer.sft_trainer import ({','.join(x)})", locals())
+ try:
+ exec(f"from trl.trainer.sft_trainer import ({','.join(x)})", locals())
+ except ImportError:
+ for _item in x:
+ try:
+ exec(f"from trl.trainer.sft_trainer import {_item}", locals())
+ except ImportError:
+ pass
exec(function, locals(), globals())
exec(
f"trl.trainer.sft_trainer.SFTTrainer.{function_name} = {function_name}",
diff --git a/unsloth/trainer.py b/unsloth/trainer.py
index cb36b8639d..65abe6801f 100644
--- a/unsloth/trainer.py
+++ b/unsloth/trainer.py
@@ -73,7 +73,7 @@ def _should_auto_padding_free(config) -> bool:
or getattr(config, "packing", False)
):
return False
- return not getattr(config, "padding_free", False)
+ return getattr(config, "padding_free", None) is None
def _disable_sample_packing(config):
@@ -392,6 +392,7 @@ def _patch_sft_trainer_auto_packing(trl_module):
packing_active = True
logger.info("Unsloth: Sample packing enabled for SFTTrainer instance.")
+ # Resolve padding_free: None (default) = auto-enable unless env-disabled or packing
auto_padding_free_active = False
padding_free_requested = getattr(config_arg, "padding_free", None) is True
if not blocked:
diff --git a/unsloth/utils/attention_dispatch.py b/unsloth/utils/attention_dispatch.py
index a7620549be..72d52ab376 100644
--- a/unsloth/utils/attention_dispatch.py
+++ b/unsloth/utils/attention_dispatch.py
@@ -20,6 +20,7 @@ from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Optional, Tuple
+import torch
from torch import Tensor
from torch.nn.functional import scaled_dot_product_attention
@@ -119,6 +120,19 @@ def run_attention(
backend = config.backend
if backend == FLASH_VARLEN and context.seq_info is None:
backend = FLASH_DENSE if HAS_FLASH_ATTENTION else SDPA
+
+ # [TODO] Flash attention does not support arbitrary attention masks (only
+ # causal via flag). When a padding mask is present (e.g. left-padded
+ # batched generation), fall back to SDPA which consumes attn_mask.
+ # xFormers also does not thread context.attention_mask through, so the
+ # same fallback applies.
+ if context.attention_mask is not None and backend in (
+ FLASH_DENSE,
+ FLASH_VARLEN,
+ XFORMERS,
+ ):
+ backend = SDPA
+
flash_dense_kwargs = config.flash_dense_kwargs or {}
flash_varlen_kwargs = config.flash_varlen_kwargs or {}
sdpa_kwargs = config.sdpa_kwargs or {}
@@ -234,14 +248,79 @@ def run_attention(
else:
q_len_local = Q.shape[-2]
k_len_local = K.shape[-2]
+ # ---- SDPA mask normalization for left padding / 2D masks ----
+ if local_mask is not None and isinstance(local_mask, torch.Tensor):
+ local_mask = local_mask.to(device = Q.device)
+
+ if local_mask.dim() == 2:
+ # key padding keep mask: (bsz, k_len), 1/True = real token
+ if local_mask.dtype == torch.bool:
+ key_keep = local_mask
+ else:
+ # tokenizer attention_mask is typically int 0/1
+ key_keep = local_mask != 0
+
+ past_len = (
+ k_len_local - q_len_local
+ ) # works for prefill (0) and decode
+ q_pos = torch.arange(
+ past_len, past_len + q_len_local, device = Q.device
+ )
+ k_pos = torch.arange(k_len_local, device = Q.device)
+
+ causal_keep = (
+ k_pos[None, :] <= q_pos[:, None]
+ ) # True = allowed (SDPA)
+ if sliding_window is not None:
+ causal_keep &= k_pos[None, :] >= (
+ q_pos[:, None] - (sliding_window - 1)
+ )
+
+ # (bsz, 1, q_len, k_len) boolean keep mask
+ local_mask = (
+ causal_keep[None, None, :, :] & key_keep[:, None, None, :]
+ )
+
+ elif local_mask.dim() == 3:
+ # (bsz, q_len, k_len) -> (bsz, 1, q_len, k_len)
+ local_mask = local_mask[:, None, :, :]
+
+ elif local_mask.dim() == 4:
+ if local_mask.dtype != torch.bool:
+ # Use boolean keep masks for better SDPA stability.
+ local_mask = local_mask.eq(0)
+ else:
+ raise ValueError(
+ f"Unsupported SDPA attention_mask rank: {local_mask.dim()}"
+ )
+
+ # Avoid NaNs from fully-masked rows (common with left padding).
+ if local_mask.dtype == torch.bool:
+ no_allowed = ~local_mask.any(
+ dim = -1, keepdim = True
+ ) # (bsz,1,q_len,1)
+ local_mask = local_mask | no_allowed
+
is_causal_local = local_mask is None and q_len_local == k_len_local
kwargs = dict(sdpa_kwargs)
kwargs.setdefault("attn_mask", local_mask)
kwargs.setdefault("is_causal", is_causal_local)
- if SDPA_HAS_GQA:
- kwargs.setdefault("enable_gqa", config.n_groups != 1)
+ use_sdpa_gqa = SDPA_HAS_GQA and config.n_groups != 1
+ if (
+ use_sdpa_gqa
+ and (not requires_grad)
+ and isinstance(local_mask, torch.Tensor)
+ and local_mask.dim() >= 3
+ and local_mask.shape[0] > 1
+ ):
+ # Batched masked inference has shown row-coupled drift with SDPA GQA.
+ # Fall back to explicit KV expansion for deterministic row-wise behavior.
+ use_sdpa_gqa = False
+
+ if use_sdpa_gqa:
+ kwargs.setdefault("enable_gqa", True)
out = scaled_dot_product_attention(Q, K, V, **kwargs)
return out.transpose(1, 2)