Merge branch 'main' into nightly

This commit is contained in:
Daniel Han 2026-02-10 07:19:28 -08:00
commit 86d7bf6c4c
22 changed files with 1296 additions and 667 deletions

View file

@ -1,6 +1,6 @@
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.14.14
rev: v0.15.0
hooks:
- id: ruff
args:

View file

@ -44,36 +44,36 @@ Notebooks are beginner friendly. Read our [guide](https://unsloth.ai/docs/get-st
pip install unsloth
```
### Windows
For Windows, `pip install unsloth` works only if you have Pytorch installed. Read our [Windows Guide](https://unsloth.ai/docs/get-started/install-and-update/windows-installation).
For Windows, `pip install unsloth` works only if you have Pytorch installed. Read our [Windows Guide](https://unsloth.ai/docs/get-started/install/windows-installation).
### 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-and-update/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/basics/fine-tuning-llms-with-blackwell-rtx-50-series-and-unsloth) and [DGX Spark Guide](https://unsloth.ai/docs/basics/fine-tuning-llms-with-nvidia-dgx-spark-and-unsloth) for more details.
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.
## 🦥 Unsloth News
- **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/new/500k-context-length-fine-tuning)
- **FP8 Reinforcement Learning**: You can now do FP8 GRPO on consumer GPUs. [Blog](https://unsloth.ai/docs/new/fp8-reinforcement-learning) • [Notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_8B_FP8_GRPO.ipynb)
- **DeepSeek-OCR**: Fine-tune to improve language understanding by 89%. [Guide](https://unsloth.ai/docs/models/deepseek-ocr-how-to-run-and-fine-tune) • [Notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Deepseek_OCR_(3B).ipynb)
- **Docker**: Use Unsloth with no setup & environment issues with our new image. [Guide](https://unsloth.ai/docs/new/how-to-fine-tune-llms-with-unsloth-and-docker) • [Docker image](https://hub.docker.com/r/unsloth/unsloth)
- **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)
- **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.
<details>
<summary>Click for more news</summary>
- **Quantization-Aware Training**: We collabed with Pytorch, recovering ~70% accuracy. [Read blog](https://unsloth.ai/docs/basics/quantization-aware-training-qat)
- **Quantization-Aware Training**: We collabed with Pytorch, recovering ~70% accuracy. [Read blog](https://unsloth.ai/docs/blog/quantization-aware-training-qat)
- **Memory-efficient RL**: We're introducing even better RL. Our new kernels & algos allows faster RL with 50% less VRAM & 10× more context. [Read blog](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/memory-efficient-rl)
- **Mistral 3**: Run Ministral 3 or Devstral 2 and fine-tune with vision/RL sodoku notebooks. [Guide](https://unsloth.ai/docs/models/ministral-3) • [Notebooks](https://unsloth.ai/docs/models/ministral-3#fine-tuning-ministral-3)
- **Mistral 3**: Run Ministral 3 or Devstral 2 and fine-tune with vision/RL sudoku notebooks. [Guide](https://unsloth.ai/docs/models/tutorials/ministral-3) • [Notebooks](https://unsloth.ai/docs/models/ministral-3#fine-tuning-ministral-3)
- **Gemma 3n** by Google: [Read Blog](https://unsloth.ai/docs/models/gemma-3-how-to-run-and-fine-tune/gemma-3n-how-to-run-and-fine-tune). We [uploaded GGUFs, 4-bit models](https://huggingface.co/collections/unsloth/gemma-3n-685d3874830e49e1c93f9339).
- **[Text-to-Speech (TTS)](https://unsloth.ai/docs/basics/text-to-speech-tts-fine-tuning)** is now supported, including `sesame/csm-1b` and STT `openai/whisper-large-v3`.
- **[Qwen3](https://unsloth.ai/docs/models/qwen3-how-to-run-and-fine-tune)** is now supported. Qwen3-30B-A3B fits on 17.5GB VRAM.
- Introducing **[Dynamic 2.0](https://unsloth.ai/docs/basics/unsloth-dynamic-2.0-ggufs)** quants that set new benchmarks on 5-shot MMLU & Aider Polyglot.
- [**EVERYTHING** is now supported](https://unsloth.ai/blog/gemma3#everything) - all models (TTS, BERT, Mamba), FFT, etc. [MultiGPU](https://unsloth.ai/docs/basics/multi-gpu-training-with-unsloth) coming soon. Enable FFT with `full_finetuning = True`, 8-bit with `load_in_8bit = True`.
- [**EVERYTHING** is now supported](https://unsloth.ai/blog/gemma3#everything) - all models (TTS, BERT, Mamba), FFT, etc. [MultiGPU](https://unsloth.ai/docs/basics/multi-gpu-training-with-unsloth) is now supported. Enable FFT with `full_finetuning = True`, 8-bit with `load_in_8bit = True`.
- 📣 [DeepSeek-R1](https://unsloth.ai/blog/deepseek-r1) - run or fine-tune them [with our guide](https://unsloth.ai/blog/deepseek-r1). All model uploads: [here](https://huggingface.co/collections/unsloth/deepseek-r1-all-versions-678e1c48f5d2fce87892ace5).
- 📣 Introducing Long-context [Reasoning (GRPO)](https://unsloth.ai/blog/grpo) in Unsloth. Train your own reasoning model with just 5GB VRAM. Transform Llama, Phi, Mistral etc. into reasoning LLMs!
- 📣 Introducing Unsloth [Dynamic 4-bit Quantization](https://unsloth.ai/blog/dynamic-4bit)! We dynamically opt not to quantize certain parameters and this greatly increases accuracy while only using <10% more VRAM than BnB 4-bit. See our collection on [Hugging Face here.](https://huggingface.co/collections/unsloth/unsloth-4-bit-dynamic-quants-67503bb873f89e15276c44e7)
@ -92,24 +92,24 @@ For RTX 50x, B200, 6000 GPUs: `pip install unsloth`. Read our [Blackwell Guide](
| <img width="15" src="https://redditinc.com/hs-fs/hubfs/Reddit%20Inc/Brand/Reddit_Logo.png" />  **r/unsloth Reddit** | [Join Reddit community](https://reddit.com/r/unsloth) |
| 📚 **Documentation & Wiki** | [Read Our Docs](https://unsloth.ai/docs) |
| <img width="13" src="https://upload.wikimedia.org/wikipedia/commons/0/09/X_(formerly_Twitter)_logo_late_2025.svg" />  **Twitter (aka X)** | [Follow us on X](https://twitter.com/unslothai) |
| 💾 **Installation** | [Pip & Docker Install](https://unsloth.ai/docs/get-started/install-and-update) |
| 💾 **Installation** | [Pip & Docker Install](https://unsloth.ai/docs/get-started/install) |
| 🔮 **Our Models** | [Unsloth Catalog](https://unsloth.ai/docs/get-started/unsloth-model-catalog) |
| ✍️ **Blog** | [Read our Blogs](https://unsloth.ai/blog) |
## ⭐ Key Features
* Supports **full-finetuning**, pretraining, 4b-bit, 16-bit and **FP8** training
* Supports **full-finetuning**, pretraining, 4-bit, 16-bit and **FP8** training
* 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-and-update/amd) and Intel GPUs. Minimum CUDA Capability 7.0 (V100, T4, Titan V, RTX 20, 30, 40x, A100, H100, L40 etc)
* 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**
* All kernels written in OpenAI's Triton language. Manual backprop engine.
* If you trained a model with 🦥Unsloth, you can use this cool sticker!   <img src="https://raw.githubusercontent.com/unslothai/unsloth/main/images/made with unsloth.png" width="200" align="center" />
## 💾 Install Unsloth
You can also see our docs for more detailed installation and updating instructions [here](https://unsloth.ai/docs/get-started/install-and-update).
You can also see our docs for more detailed installation and updating instructions [here](https://unsloth.ai/docs/get-started/install).
Unsloth supports Python 3.13 or lower.
@ -128,17 +128,17 @@ See [here](#advanced-pip-installation) for advanced pip install instructions.
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).
3. **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-and-update/windows-installation#method-3-windows-directly).
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).
5. **Install CUDA Toolkit:**
3. **Install CUDA Toolkit:**
Follow the instructions to install [CUDA Toolkit](https://developer.nvidia.com/cuda-toolkit-archive).
6. **Install PyTorch:**
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/).
7. **Install Unsloth:**
5. **Install Unsloth:**
```python
pip install unsloth
@ -163,23 +163,21 @@ pip install unsloth
```
Check if `xformers` succeeded with `python -m xformers.info` Go to https://github.com/facebookresearch/xformers. Another option is to install `flash-attn` for Ampere GPUs and ignore `xformers`
5. For GRPO runs, you can try installing `vllm` and seeing if `pip install vllm` succeeds.
6. Double check that your versions of Python, CUDA, CUDNN, `torch`, `triton`, and `xformers` are compatible with one another. The [PyTorch Compatibility Matrix](https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix) may be useful.
5. Finally, install `bitsandbytes` and check it with `python -m bitsandbytes`
4. For GRPO runs, you can try installing `vllm` and seeing if `pip install vllm` succeeds.
5. Double check that your versions of Python, CUDA, CUDNN, `torch`, `triton`, and `xformers` are compatible with one another. The [PyTorch Compatibility Matrix](https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix) may be useful.
6. Finally, install `bitsandbytes` and check it with `python -m bitsandbytes`
### Conda Installation (Optional)
`⚠Only use Conda if you have it. If not, use Pip`. Select either `pytorch-cuda=11.8,12.1` for CUDA 11.8 or CUDA 12.1. We support `python=3.10,3.11,3.12`.
`⚠Only use Conda if you have it. If not, use Pip`. We support `python=3.10,3.11,3.12,3.13`.
```bash
conda create --name unsloth_env \
python=3.11 \
pytorch-cuda=12.1 \
pytorch cudatoolkit xformers -c pytorch -c nvidia -c xformers \
-y
conda create --name unsloth_env python==3.12 -y
conda activate unsloth_env
pip install unsloth
```
Use `nvidia-smi` to get the correct CUDA version like 13.0 which becomes `cu130`
```bash
pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu130
pip3 install unsloth
```
<details>
<summary>If you're looking to install Conda in a Linux environment, <a href="https://docs.anaconda.com/miniconda/">read here</a>, or run the below 🔽</summary>
@ -269,7 +267,7 @@ print(f'pip install --upgrade pip && pip install --no-deps git+https://github.co
```
### Docker Installation
You can use our pre-built Docker container with all dependencies to use Unsloth instantly with no setup required.
[Read our guide](https://unsloth.ai/docs/get-started/install-and-update/docker).
[Read our guide](https://unsloth.ai/docs/get-started/install/docker).
This container requires installing [NVIDIA's Container Toolkit](https://docs.nvidia.com/datacenter/cloud-native/container-toolkit/latest/install-guide.html).
@ -285,13 +283,13 @@ Access Jupyter Lab at `http://localhost:8888` and start fine-tuning!
## 📜 Documentation
* Go to our official [Documentation](https://unsloth.ai/docs) for [running models](https://unsloth.ai/docs/basics/inference-and-deployment), [saving to GGUF](https://unsloth.ai/docs/basics/inference-and-deployment/saving-to-gguf), [checkpointing](https://unsloth.ai/docs/basics/finetuning-from-last-checkpoint), [evaluation](https://unsloth.ai/docs/get-started/fine-tuning-llms-guide#evaluation) and more!
* Read our Guides for: [Fine-tuning](https://unsloth.ai/docs/get-started/fine-tuning-llms-guide), [Reinforcement Learning](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide), [Text-to-Speech (TTS)](https://unsloth.ai/docs/basics/text-to-speech-tts-fine-tuning), [Vision](https://unsloth.ai/docs/basics/vision-fine-tuning) and [any model](https://unsloth.ai/docs/models/tutorials-how-to-fine-tune-and-run-llms).
* Read our Guides for: [Fine-tuning](https://unsloth.ai/docs/get-started/fine-tuning-llms-guide), [Reinforcement Learning](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide), [Text-to-Speech (TTS)](https://unsloth.ai/docs/basics/text-to-speech-tts-fine-tuning), [Vision](https://unsloth.ai/docs/basics/vision-fine-tuning) and [any model](https://unsloth.ai/docs/models/tutorials).
* We support Huggingface's transformers, TRL, Trainer, Seq2SeqTrainer and Pytorch code.
Unsloth example code to fine-tune gpt-oss-20b:
```python
from unsloth import FastLanguageModel, FastModel
from unsloth import FastLanguageModel, FastModel, FastVisionModel
import torch
from trl import SFTTrainer, SFTConfig
from datasets import load_dataset
@ -306,9 +304,9 @@ fourbit_models = [
] # More models at https://huggingface.co/unsloth
model, tokenizer = FastModel.from_pretrained(
model, tokenizer = FastLanguageModel.from_pretrained(
model_name = "unsloth/gpt-oss-20b",
max_seq_length = 2048, # Choose any for long context!
max_seq_length = max_seq_length, # Choose any for long context!
load_in_4bit = True, # 4-bit quantization. False = 16-bit LoRA.
load_in_8bit = False, # 8-bit quantization
load_in_16bit = False, # 16-bit LoRA
@ -361,14 +359,14 @@ trainer.train()
<a name="RL"></a>
## 💡 Reinforcement Learning
[RL](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide) including [GRPO](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide#training-with-grpo), [GSPO](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/gspo-reinforcement-learning), [**FP8** training](https://unsloth.ai/docs/new/fp8-reinforcement-learning), DrGRPO, DAPO, PPO, Reward Modelling, Online DPO all work with Unsloth.
[RL](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide) including [GRPO](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide#training-with-grpo), [GSPO](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/gspo-reinforcement-learning), [**FP8** training](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/fp8-reinforcement-learning), DrGRPO, DAPO, PPO, Reward Modelling, Online DPO all work with Unsloth.
Read our [Reinforcement Learning Guide](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide) or our [advanced RL docs](https://unsloth.ai/docs/get-started/reinforcement-learning-rl-guide/advanced-rl-documentation) for batching, generation & training parameters.
List of RL notebooks:
- gpt-oss GSPO notebook: [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/gpt-oss-(20B)-GRPO.ipynb)
- - ***FP8*** Qwen3-8B GRPO notebook (L4): [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_8B_FP8_GRPO.ipynb)
- Qwen2.3-VL GSPO notebook: [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_VL_(8B)-Vision-GRPO.ipynb)
- gpt-oss GRPO notebook: [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/gpt-oss-(20B)-GRPO.ipynb)
- ***FP8*** Qwen3-8B GRPO notebook (L4): [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_8B_FP8_GRPO.ipynb)
- Qwen3-VL GSPO notebook: [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_VL_(8B)-Vision-GRPO.ipynb)
- Advanced Qwen3 GRPO notebook: [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_(4B)-GRPO.ipynb)
- ORPO notebook: [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3_(8B)-ORPO.ipynb)
- DPO Zephyr notebook: [Link](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Zephyr_(7B)-DPO.ipynb)
@ -420,7 +418,7 @@ You can cite the Unsloth repo as follows:
@software{unsloth,
author = {Daniel Han, Michael Han and Unsloth team},
title = {Unsloth},
url = {http://github.com/unslothai/unsloth},
url = {https://github.com/unslothai/unsloth},
year = {2023}
}
```

View file

@ -178,13 +178,29 @@ class _DummyModel(torch.nn.Module):
class _DummyTrainer:
def __init__(self):
self.args = SimpleNamespace(remove_unused_columns = True)
self.data_collator = DataCollatorForLanguageModeling(
pad_token_id = 0,
completion_only_loss = False,
padding_free = True,
return_position_ids = False,
return_tensors = "pt",
)
collator_args = {
"pad_token_id": 0,
"completion_only_loss": False,
"return_tensors": "pt",
}
optional_flags = [
{"padding_free": True, "return_position_ids": False},
{"padding_free": True},
{},
]
for extra in optional_flags:
try:
self.data_collator = DataCollatorForLanguageModeling(
**collator_args, **extra
)
break
except TypeError:
continue
# Ensure attributes exist even if the constructor did not accept them
if not hasattr(self.data_collator, "padding_free"):
self.data_collator.padding_free = True
if not hasattr(self.data_collator, "return_position_ids"):
self.data_collator.return_position_ids = False
class _PaddingFreeCollator:

View file

@ -128,6 +128,7 @@ from .import_fixes import (
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,
ignore_logger_messages,
patch_ipykernel_hf_xet,
@ -139,6 +140,7 @@ from .import_fixes import (
fix_executorch,
patch_vllm_for_notebooks,
patch_torchcodec_audio_decoder,
disable_torchcodec_if_broken,
)
fix_xformers_performance_issue()
@ -147,6 +149,7 @@ fix_vllm_aimv2_issue()
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()
ignore_logger_messages()
patch_ipykernel_hf_xet()
@ -158,12 +161,14 @@ patch_openspiel_env_async()
fix_executorch()
patch_vllm_for_notebooks()
patch_torchcodec_audio_decoder()
disable_torchcodec_if_broken()
del fix_xformers_performance_issue
del fix_vllm_aimv2_issue
del check_vllm_torch_sm100_compatibility
del fix_vllm_guided_decoding_params
del fix_vllm_pdl_blackwell
del fix_triton_compiled_kernel_missing_attrs
del fix_rocm_triton_key_error
del ignore_logger_messages
del patch_ipykernel_hf_xet
@ -175,6 +180,7 @@ del patch_openspiel_env_async
del fix_executorch
del patch_vllm_for_notebooks
del patch_torchcodec_audio_decoder
del disable_torchcodec_if_broken
# Torch 2.4 has including_emulation
if DEVICE_TYPE == "cuda":

View file

@ -36,6 +36,7 @@ import shutil
from .tokenizer_utils import *
from .models._utils import patch_tokenizer
import re
from .ollama_template_mappers import OLLAMA_TEMPLATES
from unsloth_zoo.dataset_utils import (
train_on_responses_only,
standardize_data_formats,
@ -43,6 +44,8 @@ from unsloth_zoo.dataset_utils import (
standardize_sharegpt = standardize_data_formats
CHAT_TEMPLATES = {}
DEFAULT_SYSTEM_MESSAGE = {}
def _ollama_template(name: str):
return OLLAMA_TEMPLATES[name]
# =========================================== Unsloth
# Unsloth efficient template leverages from Zephyr
@ -68,18 +71,7 @@ unsloth_template = \
"{{ '>>> Assistant: ' }}"\
"{% endif %}"
unsloth_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}{{ .System }}
{{ end }}{{ if .Prompt }}>>> User: {{ .Prompt }}
{{ end }}>>> Assistant: {{ .Response }}{__EOS_TOKEN__}
"""
PARAMETER stop "{__EOS_TOKEN__}"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
SYSTEM """You are a helpful assistant to the user"""
'''
unsloth_ollama = _ollama_template("unsloth")
unsloth_eos_token = "eos_token"
CHAT_TEMPLATES["unsloth"] = (unsloth_template, unsloth_eos_token, False, unsloth_ollama,)
@ -101,20 +93,7 @@ zephyr_template = \
"{{ '<|assistant|>\n' }}"\
"{% endif %}"
zephyr_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}<|system|>
{{ .System }}{__EOS_TOKEN__}
{{ end }}{{ if .Prompt }}<|user|>
{{ .Prompt }}{__EOS_TOKEN__}
{{ end }}<|assistant|>
{{ .Response }}{__EOS_TOKEN__}
"""
PARAMETER stop "{__EOS_TOKEN__}"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
zephyr_ollama = _ollama_template("zephyr")
zephyr_eos_token = "eos_token"
CHAT_TEMPLATES["zephyr"] = (zephyr_template, zephyr_eos_token, False, zephyr_ollama,)
@ -136,21 +115,7 @@ chatml_template = \
"{{ '<|im_start|>assistant\n' }}"\
"{% endif %}"
chatml_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}<|im_start|>system
{{ .System }}<|im_end|>
{{ end }}{{ if .Prompt }}<|im_start|>user
{{ .Prompt }}<|im_end|>
{{ end }}<|im_start|>assistant
{{ .Response }}<|im_end|>
"""
PARAMETER stop "<|im_start|>"
PARAMETER stop "<|im_end|>"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
chatml_ollama = _ollama_template("chatml")
chatml_eos_token = "<|im_end|>"
CHAT_TEMPLATES["chatml"] = (chatml_template, chatml_eos_token, True, chatml_ollama,)
@ -182,14 +147,7 @@ mistral_template = \
"{% endfor %}"
# Ollama from https://www.ollama.com/library/mistral
mistral_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """[INST] {{ if .System }}{{ .System }} {{ end }}{{ .Prompt }} [/INST]"""
PARAMETER stop "{__EOS_TOKEN__}"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
mistral_ollama = _ollama_template("mistral")
mistral_eos_token = "eos_token"
CHAT_TEMPLATES["mistral"] = (mistral_template, mistral_eos_token, False, mistral_ollama,)
@ -220,16 +178,7 @@ llama_template = \
"{% endfor %}"
# Ollama from https://www.ollama.com/library/llama3
llama_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """[INST] <<SYS>>{{ .System }}<</SYS>>
{{ .Prompt }} [/INST]"""
PARAMETER stop "{__EOS_TOKEN__}"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
llama_ollama = _ollama_template("llama")
llama_eos_token = "eos_token"
CHAT_TEMPLATES["llama"] = (llama_template, llama_eos_token, False, llama_ollama,)
@ -260,14 +209,7 @@ vicuna_template = \
"{% endif %}"
# Ollama from https://www.ollama.com/library/vicuna
vicuna_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}{{ .System }} {{ end }}{{ if .Prompt }}USER: {{ .Prompt }} {{ end }}ASSISTANT: {{ .Response }} {__EOS_TOKEN__}"""
PARAMETER stop "{__EOS_TOKEN__}"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
vicuna_ollama = _ollama_template("vicuna")
vicuna_eos_token = "eos_token"
CHAT_TEMPLATES["vicuna"] = (vicuna_template, vicuna_eos_token, False, vicuna_ollama,)
@ -297,18 +239,7 @@ vicuna_old_template = \
"{{ '### Assistant:' }}"\
"{% endif %}"
vicuna_old_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}{{ .System }}
{{ end }}{{ if .Prompt }}### Human: {{ .Prompt }}
{{ end }}### Assistant: {{ .Response }}{__EOS_TOKEN__}
"""
PARAMETER stop "{__EOS_TOKEN__}"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
SYSTEM """A chat between a curious human and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the human's questions."""
'''
vicuna_old_ollama = _ollama_template("vicuna_old")
vicuna_old_eos_token = "eos_token"
CHAT_TEMPLATES["vicuna_old"] = (vicuna_old_template, vicuna_old_eos_token, False, vicuna_old_ollama,)
@ -341,23 +272,7 @@ alpaca_template = \
"{{ '### Response:\n' }}"\
"{% endif %}"
alpaca_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}{{ .System }}
{{ end }}{{ if .Prompt }}### Instruction:
{{ .Prompt }}{{ end }}
### Response:
{{ .Response }}{__EOS_TOKEN__}
"""
PARAMETER stop "{__EOS_TOKEN__}"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
SYSTEM """Below are some instructions that describe some tasks. Write responses that appropriately complete each request."""
'''
alpaca_ollama = _ollama_template("alpaca")
alpaca_eos_token = "eos_token"
CHAT_TEMPLATES["alpaca"] = (alpaca_template, alpaca_eos_token, False, alpaca_ollama,)
@ -387,21 +302,7 @@ gemma_template = \
"{% endif %}"
# Ollama from https://www.ollama.com/library/gemma
gemma_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """<start_of_turn>user
{{ if .System }}{{ .System }} {{ end }}{{ .Prompt }}<end_of_turn>
<start_of_turn>model
{{ .Response }}<end_of_turn>
"""
PARAMETER repeat_penalty 1
PARAMETER stop "<start_of_turn>"
PARAMETER stop "<end_of_turn>"
PARAMETER penalize_newline false
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
gemma_ollama = _ollama_template("gemma")
gemma_eos_token = "<end_of_turn>"
CHAT_TEMPLATES["gemma"] = (gemma_template, gemma_eos_token, True, gemma_ollama,)
@ -411,23 +312,7 @@ DEFAULT_SYSTEM_MESSAGE["gemma"] = None # No system message in Gemma
# We find using <eos> is still more appropriate!
gemma_chatml_template = "{{ bos_token }}" + chatml_template
gemma_chatml_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}<|im_start|>system
{{ .System }}<|im_end|>
{{ end }}{{ if .Prompt }}<|im_start|>user
{{ .Prompt }}<|im_end|>
{{ end }}<|im_start|>assistant
{{ .Response }}<|im_end|>
"""
PARAMETER repeat_penalty 1
PARAMETER stop "<|im_start|>"
PARAMETER stop "<|im_end|>"
PARAMETER penalize_newline false
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
gemma_chatml_ollama = _ollama_template("gemma_chatml")
gemma_chatml_eos_token = (
{"<start_of_turn>" : "<|im_start|>", "<eos>" : "<|im_end|>"},
@ -440,14 +325,14 @@ DEFAULT_SYSTEM_MESSAGE["gemma_chatml"] = None # No system message in Gemma
# Same as Gemma 1, but with sliding window attention!
# https://ollama.com/library/gemma2/blobs/6522ca797f47
gemma2_template = gemma_template
gemma2_ollama = gemma_ollama + "PARAMETER num_ctx 4096\n"
gemma2_ollama = _ollama_template("gemma2")
gemma2_eos_token = "<end_of_turn>"
CHAT_TEMPLATES["gemma2"] = (gemma2_template, gemma2_eos_token, True, gemma2_ollama,)
DEFAULT_SYSTEM_MESSAGE["gemma2"] = None # No system message in Gemma 2
# =========================================== Gemma 2 with ChatML instead
gemma2_chatml_template = gemma_chatml_template
gemma2_chatml_ollama = gemma_chatml_ollama + "PARAMETER num_ctx 4096\n"
gemma2_chatml_ollama = _ollama_template("gemma2_chatml")
gemma2_chatml_eos_token = gemma_chatml_eos_token
CHAT_TEMPLATES["gemma2_chatml"] = (gemma2_chatml_template, gemma2_chatml_eos_token, True, gemma2_chatml_ollama,)
DEFAULT_SYSTEM_MESSAGE["gemma2_chatml"] = None # No system message in Gemma 2
@ -470,22 +355,7 @@ llama3_template = \
"{% endif %}"
# Ollama from https://www.ollama.com/library/llama3
llama3_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}<|start_header_id|>system<|end_header_id|>
{{ .System }}<|eot_id|>{{ end }}{{ if .Prompt }}<|start_header_id|>user<|end_header_id|>
{{ .Prompt }}<|eot_id|>{{ end }}<|start_header_id|>assistant<|end_header_id|>
{{ .Response }}<|eot_id|>"""
PARAMETER stop "<|start_header_id|>"
PARAMETER stop "<|end_header_id|>"
PARAMETER stop "<|eot_id|>"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
llama3_ollama = _ollama_template("llama-3")
llama3_template_eos_token = "eos_token"
@ -513,22 +383,7 @@ phi3_template = \
"{% endif %}"
# Ollama from https://www.ollama.com/library/phi3
phi3_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}<|system|>
{{ .System }}<|end|>
{{ end }}{{ if .Prompt }}<|user|>
{{ .Prompt }}<|end|>
{{ end }}<|assistant|>
{{ .Response }}<|end|>
"""
PARAMETER stop "<|end|>"
PARAMETER stop "<|user|>"
PARAMETER stop "<|assistant|>"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
phi3_ollama = _ollama_template("phi-3")
phi3_template_eos_token = "<|end|>"
CHAT_TEMPLATES["phi-3"] = (phi3_template, phi3_template_eos_token, False, phi3_ollama,)
@ -670,65 +525,7 @@ llama31_template = \
"""
# Ollama from https://ollama.com/library/llama3.1 (needs updating!)
llama31_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .Messages }}
{{- if or .System .Tools }}<|start_header_id|>system<|end_header_id|>
{{- if .System }}
{{ .System }}
{{- end }}
{{- if .Tools }}
You are a helpful assistant with tool calling capabilities. When you receive a tool call response, use the output to format an answer to the original use question.
{{- end }}
{{- end }}<|eot_id|>
{{- range $i, $_ := .Messages }}
{{- $last := eq (len (slice $.Messages $i)) 1 }}
{{- if eq .Role "user" }}<|start_header_id|>user<|end_header_id|>
{{- if and $.Tools $last }}
Given the following functions, please respond with a JSON for a function call with its proper arguments that best answers the given prompt.
Respond in the format {"name": function name, "parameters": dictionary of argument name and its value}. Do not use variables.
{{ $.Tools }}
{{- end }}
{{ .Content }}<|eot_id|>{{ if $last }}<|start_header_id|>assistant<|end_header_id|>
{{ end }}
{{- else if eq .Role "assistant" }}<|start_header_id|>assistant<|end_header_id|>
{{- if .ToolCalls }}
{{- range .ToolCalls }}{"name": "{{ .Function.Name }}", "parameters": {{ .Function.Arguments }}}{{ end }}
{{- else }}
{{ .Content }}{{ if not $last }}<|eot_id|>{{ end }}
{{- end }}
{{- else if eq .Role "tool" }}<|start_header_id|>ipython<|end_header_id|>
{{ .Content }}<|eot_id|>{{ if $last }}<|start_header_id|>assistant<|end_header_id|>
{{ end }}
{{- end }}
{{- end }}
{{- else }}
{{- if .System }}<|start_header_id|>system<|end_header_id|>
{{ .System }}<|eot_id|>{{ end }}{{ if .Prompt }}<|start_header_id|>user<|end_header_id|>
{{ .Prompt }}<|eot_id|>{{ end }}<|start_header_id|>assistant<|end_header_id|>
{{ end }}{{ .Response }}{{ if .Response }}<|eot_id|>{{ end }}"""
PARAMETER stop "<|start_header_id|>"
PARAMETER stop "<|end_header_id|>"
PARAMETER stop "<|eot_id|>"
PARAMETER stop "<|eom_id|>"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
llama31_ollama = _ollama_template("llama-3.1")
llama31_template_eos_token = "eos_token"
CHAT_TEMPLATES["llama-3.1"] = (llama31_template, llama31_template_eos_token, False, llama31_ollama,)
@ -796,64 +593,7 @@ qwen25_template = \
# Ollama from https://ollama.com/library/qwen2.5/blobs/eb4402837c78
qwen25_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{- if .Messages }}
{{- if or .System .Tools }}<|im_start|>system
{{- if .System }}
{{ .System }}
{{- end }}
{{- if .Tools }}
# Tools
You may call one or more functions to assist with the user query.
You are provided with function signatures within <tools></tools> XML tags:
<tools>
{{- range .Tools }}
{"type": "function", "function": {{ .Function }}}
{{- end }}
</tools>
For each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:
<tool_call>
{"name": <function-name>, "arguments": <args-json-object>}
</tool_call>
{{- end }}<|im_end|>
{{ end }}
{{- range $i, $_ := .Messages }}
{{- $last := eq (len (slice $.Messages $i)) 1 -}}
{{- if eq .Role "user" }}<|im_start|>user
{{ .Content }}<|im_end|>
{{ else if eq .Role "assistant" }}<|im_start|>assistant
{{ if .Content }}{{ .Content }}
{{- else if .ToolCalls }}<tool_call>
{{ range .ToolCalls }}{"name": "{{ .Function.Name }}", "arguments": {{ .Function.Arguments }}}
{{ end }}</tool_call>
{{- end }}{{ if not $last }}<|im_end|>
{{ end }}
{{- else if eq .Role "tool" }}<|im_start|>user
<tool_response>
{{ .Content }}
</tool_response><|im_end|>
{{ end }}
{{- if and (ne .Role "assistant") $last }}<|im_start|>assistant
{{ end }}
{{- end }}
{{- else }}
{{- if .System }}<|im_start|>system
{{ .System }}<|im_end|>
{{ end }}{{ if .Prompt }}<|im_start|>user
{{ .Prompt }}<|im_end|>
{{ end }}<|im_start|>assistant
{{ end }}{{ .Response }}{{ if .Response }}<|im_end|>{{ end }}"""
PARAMETER stop "<|im_end|>"
PARAMETER stop "<|endoftext|>"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
qwen25_ollama = _ollama_template("qwen-2.5")
qwen25_template_eos_token = "eos_token"
qwen25_default_system_message = "You are Qwen, created by Alibaba Cloud. You are a helpful assistant."
@ -891,16 +631,7 @@ _phi4_ollama_template = \
"<|im_start|><|assistant|><|im_sep|>{{ .Response }}<|im_end|>"
# Ollama from https://www.ollama.com/library/phi4 is different
phi4_ollama = \
f'''
FROM {{__FILE_LOCATION__}}
TEMPLATE """{_phi4_ollama_template}"""
PARAMETER stop "<|im_end|>"
PARAMETER stop "<|im_start|>"
PARAMETER stop "<|im_sep|>"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
phi4_ollama = _ollama_template("phi-4")
phi4_template_eos_token = "<|im_end|>"
CHAT_TEMPLATES["phi-4"] = (phi4_template, phi4_template_eos_token, False, phi4_ollama,)
@ -954,28 +685,7 @@ gemma3_template = \
"""
# Ollama from https://ollama.com/library/gemma3/blobs/e0a42594d802
gemma3_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{- range $i, $_ := .Messages }}
{{- $last := eq (len (slice $.Messages $i)) 1 }}
{{- if or (eq .Role "user") (eq .Role "system") }}<start_of_turn>user
{{ .Content }}<end_of_turn>
{{ if $last }}<start_of_turn>model
{{ end }}
{{- else if eq .Role "assistant" }}<start_of_turn>model
{{ .Content }}{{ if not $last }}<end_of_turn>
{{ end }}
{{- end }}
{{- end }}"""
PARAMETER stop "<end_of_turn>"
PARAMETER stop "<eos>"
PARAMETER temperature 0.1
PARAMETER min_p 0.0
PARAMETER top_k 64
PARAMETER top_p 0.95
PARAMETER num_predict 32768
'''
gemma3_ollama = _ollama_template("gemma-3")
gemma3_template_eos_token = "<end_of_turn>"
CHAT_TEMPLATES["gemma-3"] = (gemma3_template, gemma3_template_eos_token, False, gemma3_ollama,)
@ -1088,69 +798,7 @@ qwen3_template = \
{%- endif %}
"""
# Ollama template for Qwen-3 (see https://ollama.com/library/qwen3/blobs/eb4402837c78)
qwen3_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{- if .Messages }}
{{- if or .System .Tools }}<|im_start|>system
{{- if .System }}
{{ .System }}
{{- end }}
{{- if .Tools }}
# Tools
You may call one or more functions to assist with the user query.
You are provided with function signatures within <tools></tools> XML tags:
<tools>
{{- range .Tools }}
{"type": "function", "function": {{ .Function }}}
{{- end }}
</tools>
For each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:
<tool_call>
{"name": <function-name>, "arguments": <args-json-object>}
</tool_call>
{{- end }}<|im_end|>
{{ end }}
{{- range $i, $_ := .Messages }}
{{- $last := eq (len (slice $.Messages $i)) 1 -}}
{{- if eq .Role "user" }}<|im_start|>user
{{ .Content }}<|im_end|>
{{ else if eq .Role "assistant" }}<|im_start|>assistant
{{ if .Content }}{{ .Content }}
{{- else if .ToolCalls }}<tool_call>
{{ range .ToolCalls }}{"name": "{{ .Function.Name }}", "arguments": {{ .Function.Arguments }}}
{{ end }}</tool_call>
{{- end }}{{ if not $last }}<|im_end|>
{{ end }}
{{- else if eq .Role "tool" }}<|im_start|>user
<tool_response>
{{ .Content }}
</tool_response><|im_end|>
{{ end }}
{{- if and (ne .Role "assistant") $last }}<|im_start|>assistant
{{ end }}
{{- end }}
{{- else }}
{{- if .System }}<|im_start|>system
{{ .System }}<|im_end|>
{{ end }}{{ if .Prompt }}<|im_start|>user
{{ .Prompt }}<|im_end|>
{{ end }}<|im_start|>assistant
{{ end }}{{ .Response }}{{ if .Response }}<|im_end|>{{ end }}"""
PARAMETER stop "<|im_end|>"
PARAMETER stop "<|im_start|>"
PARAMETER temperature 0.6
PARAMETER min_p 0.0
PARAMETER top_k 20
PARAMETER top_p 0.95
PARAMETER repeat_penalty 1
'''
qwen3_ollama = _ollama_template("qwen-3")
qwen3_template_eos_token = "<|im_end|>"
CHAT_TEMPLATES["qwen-3"] = (qwen3_template, qwen3_template_eos_token, False, qwen3_ollama,)
DEFAULT_SYSTEM_MESSAGE["qwen-3"] = None # No default system message for Qwen-3
@ -1207,22 +855,7 @@ gemma3n_template = \
"""
# Ollama from https://ollama.com/library/gemma3n/blobs/e0a42594d802
gemma3n_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{- range $i, $_ := .Messages }}
{{- $last := eq (len (slice $.Messages $i)) 1 }}
{{- if or (eq .Role "user") (eq .Role "system") }}<start_of_turn>user
{{ .Content }}<end_of_turn>
{{ if $last }}<start_of_turn>model
{{ end }}
{{- else if eq .Role "assistant" }}<start_of_turn>model
{{ .Content }}{{ if not $last }}<end_of_turn>
{{ end }}
{{- end }}
{{- end }}"""
'''
gemma3n_ollama = _ollama_template("gemma-3n")
gemma3n_template_eos_token = "<end_of_turn>"
CHAT_TEMPLATES["gemma-3n"] = (gemma3n_template, gemma3n_template_eos_token, False, gemma3n_ollama,)
DEFAULT_SYSTEM_MESSAGE["gemma-3n"] = None # No system message in Gemma-3n
@ -1583,7 +1216,7 @@ gptoss_template = \
<|start|>assistant
{%- endif -%}"""
# Ollama from https://ollama.com/library/gemma3n/blobs/e0a42594d802
# Ollama from https://ollama.com/library/gpt-oss
gptoss_ollama = \
'''
FROM {__FILE_LOCATION__}
@ -1861,66 +1494,8 @@ qwen3_instruct_template = \
{{- '<|im_start|>assistant\\n' }}
{%- endif %}'''
# Ollama from https://ollama.com/library/qwen3/blobs/53e4ea15e8f5
qwen3_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """
{{- $lastUserIdx := -1 -}}
{{- range $idx, $msg := .Messages -}}
{{- if eq $msg.Role "user" }}{{ $lastUserIdx = $idx }}{{ end -}}
{{- end }}
{{- if or .System .Tools }}<|im_start|>system
{{ if .System }}
{{ .System }}
{{- end }}
{{- if .Tools }}
# Tools
You may call one or more functions to assist with the user query.
You are provided with function signatures within <tools></tools> XML tags:
<tools>
{{- range .Tools }}
{"type": "function", "function": {{ .Function }}}
{{- end }}
</tools>
For each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:
<tool_call>
{"name": <function-name>, "arguments": <args-json-object>}
</tool_call>
{{- end -}}
<|im_end|>
{{ end }}
{{- range $i, $_ := .Messages }}
{{- $last := eq (len (slice $.Messages $i)) 1 -}}
{{- if eq .Role "user" }}<|im_start|>user
{{ .Content }}<|im_end|>
{{ else if eq .Role "assistant" }}<|im_start|>assistant
{{ if (and $.IsThinkSet (and .Thinking (or $last (gt $i $lastUserIdx)))) -}}
<think>{{ .Thinking }}</think>
{{ end -}}
{{ if .Content }}{{ .Content }}
{{- else if .ToolCalls }}<tool_call>
{{ range .ToolCalls }}{"name": "{{ .Function.Name }}", "arguments": {{ .Function.Arguments }}}
{{ end }}</tool_call>
{{- end }}{{ if not $last }}<|im_end|>
{{ end }}
{{- else if eq .Role "tool" }}<|im_start|>user
<tool_response>
{{ .Content }}
</tool_response><|im_end|>
{{ end }}
{{- if and (ne .Role "assistant") $last }}<|im_start|>assistant
{{ end }}
{{- end }}
"""
'''
qwen3_template_eos_token = "<|im_end|>"
CHAT_TEMPLATES["qwen3-instruct"] = (qwen3_instruct_template, qwen3_template_eos_token, False, qwen3_ollama,)
CHAT_TEMPLATES["qwen3-instruct"] = (qwen3_instruct_template, qwen3_template_eos_token, False, _ollama_template("qwen3-instruct"),)
DEFAULT_SYSTEM_MESSAGE["qwen3-instruct"] = None # No system message in Qwen3
@ -2013,7 +1588,12 @@ qwen3_thinking_template = \
{{- '<|im_start|>assistant\n<think>\n' }}
{%- endif %}'''
CHAT_TEMPLATES["qwen3-thinking"] = (qwen3_thinking_template, qwen3_template_eos_token, False, qwen3_ollama,)
CHAT_TEMPLATES["qwen3-thinking"] = (
qwen3_thinking_template,
qwen3_template_eos_token,
False,
_ollama_template("qwen3-thinking"),
)
DEFAULT_SYSTEM_MESSAGE["qwen3-thinking"] = None # No system message in Qwen3
@ -2042,19 +1622,7 @@ starling_template = \
{%- endif %}"""
# Ollama from https://ollama.com/library/starling-lm:7b/blobs/4b21bfc435b4
starling_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}GPT4 Correct System: {{ .System }}<|end_of_turn|>
{{ end }}{{ if .Prompt }}GPT4 Correct User: {{ .Prompt }}<|end_of_turn|>
{{ end }}GPT4 Correct Assistant: {{ .Response }}<|end_of_turn|>"""
PARAMETER stop "<|end_of_turn|>"
PARAMETER stop "GPT4 Correct User:"
PARAMETER stop "GPT4 Correct Assistant:"
PARAMETER stop "GPT4 Correct System:"
PARAMETER temperature 1.5
PARAMETER min_p 0.1
'''
starling_ollama = _ollama_template("starling")
starling_template_eos_token = "<|end_of_turn|>"
CHAT_TEMPLATES["starling"] = (starling_template, starling_template_eos_token, False, starling_ollama)
@ -2072,16 +1640,7 @@ yi_chat_template = \
"""
# Ollama from https://ollama.com/library/yi:34b-chat/blobs/62fbfd9ed093
yi_chat_ollama = \
'''
FROM {__FILE_LOCATION__}
TEMPLATE """{{ if .System }}<|im_start|>system
{{ .System }}<|im_end|>
{{ end }}{{ if .Prompt }}<|im_start|>user
{{ .Prompt }}<|im_end|>
{{ end }}<|im_start|>assistant
{{ .Response }}<|im_end|>"""
'''
yi_chat_ollama = _ollama_template("yi-chat")
yi_chat_template_eos_token = "<|endoftext|>"
CHAT_TEMPLATES["yi-chat"] = (yi_chat_template, yi_chat_template_eos_token, False, yi_chat_ollama)

View file

@ -164,6 +164,39 @@ if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1":
"ignore", message = r"unclosed file.*dev/null", category = ResourceWarning
)
# torch 2.9+ pin_memory/is_pinned device arg deprecation
warnings.filterwarnings(
"ignore",
message = r"The `device` argument is deprecated",
category = DeprecationWarning,
)
warnings.filterwarnings(
"ignore",
message = r".*pin_memory.*device.*deprecated",
category = DeprecationWarning,
)
warnings.filterwarnings(
"ignore",
message = r".*is_pinned.*device.*deprecated",
category = DeprecationWarning,
)
# vllm "Level is deprecated" stderr noise
sys.stderr.add_filter("Level is deprecated")
# PydanticSerializationUnexpectedValue warning
warnings.filterwarnings(
"ignore",
message = r".*PydanticSerializationUnexpectedValue",
)
warnings.filterwarnings(
"ignore",
message = r"Expected.*but got.*with value.*is not.*subclass",
)
# Triton "df: No such file or directory" stderr noise
sys.stderr.add_filter("df: No such file")
# Fix up AttributeError: 'MessageFactory' object has no attribute 'GetPrototype'
# MUST do this at the start primarily due to tensorflow causing issues
@ -799,6 +832,54 @@ def fix_huggingface_hub():
)
def fix_triton_compiled_kernel_missing_attrs():
"""
Triton 3.6.0+ removed direct `num_ctas` and `cluster_dims` attributes from
CompiledKernel, but torch 2.9.x Inductor still expects them in
torch/_inductor/runtime/triton_heuristics.py make_launcher() (line ~1757).
The scope dict eagerly evaluates:
binary.metadata.num_ctas, *binary.metadata.cluster_dims
when hasattr(binary, "metadata") is True, but metadata lacks cluster_dims.
This crashes before reaching the new launch path that doesn't need cta_args.
Upstream fix: pytorch/pytorch@97bd4db added hasattr guards.
We monkey-patch CompiledKernel.__init__ to inject the missing attributes
so the older hasattr(binary, "num_ctas") branch succeeds instead.
"""
try:
import torch
except (ImportError, ModuleNotFoundError):
return
try:
import triton
import triton.compiler.compiler as triton_compiler
except (ImportError, ModuleNotFoundError):
return
# Only needed when the CompiledKernel class lacks num_ctas as a direct attr
# but has metadata (triton >= 3.6.0 with torch < 2.10)
_ck_cls = triton_compiler.CompiledKernel
if hasattr(_ck_cls, "num_ctas"):
return # Old triton with direct attrs -- no patch needed
_orig_init = _ck_cls.__init__
def _patched_init(self, *args, **kwargs):
_orig_init(self, *args, **kwargs)
if not hasattr(self, "num_ctas"):
self.num_ctas = getattr(self.metadata, "num_ctas", 1)
if not hasattr(self, "cluster_dims") and not hasattr(self, "clusterDims"):
self.cluster_dims = (1, 1, 1)
_ck_cls.__init__ = _patched_init
logger.info(
"Unsloth: Patched triton CompiledKernel with num_ctas/cluster_dims "
"for torch.compile compatibility."
)
def fix_rocm_triton_key_error():
"""
ROCm + torch.compile can fail if Triton lacks `triton_key`.
@ -1051,5 +1132,34 @@ def patch_torchcodec_audio_decoder():
from unsloth_zoo.dataset_utils import patch_torchcodec_audio_decoder as _patch
_patch()
except (ImportError, AttributeError):
except (ImportError, AttributeError, RuntimeError):
pass
def disable_torchcodec_if_broken():
"""Disable torchcodec in transformers if it cannot actually load.
transformers checks if torchcodec is installed via importlib.util.find_spec(),
but this returns True even when torchcodec cannot load its native libraries
(e.g., when FFmpeg is missing). This causes runtime errors when transformers
tries to use torchcodec for audio loading.
This function tests if torchcodec can actually load and if not, patches
transformers to think torchcodec is unavailable so it falls back to librosa.
"""
try:
import importlib.util
if importlib.util.find_spec("torchcodec") is None:
return # torchcodec not installed, nothing to do
# Test if torchcodec can actually load
from torchcodec.decoders import AudioDecoder
except (ImportError, RuntimeError, OSError):
# torchcodec cannot load - disable it in transformers
try:
import transformers.utils.import_utils as tf_import_utils
tf_import_utils._torchcodec_available = False
except (ImportError, AttributeError):
pass

View file

@ -103,6 +103,7 @@ from ..device_type import (
DEVICE_COUNT,
ALLOW_PREQUANTIZED_MODELS,
)
from ..import_fixes import UNSLOTH_ENABLE_LOGGING
from unsloth_zoo.log import logger
from unsloth_zoo.tokenizer_utils import (
patch_tokenizer as _patch_tokenizer,
@ -255,8 +256,45 @@ class HideLoggingMessage(logging.Filter):
return not (self.text in x.getMessage())
# Replace warning messages (analogous to HideLoggingMessage but for warnings.warn)
class ReplaceWarningMessage:
"""
Intercepts warnings.warn calls and replaces matching messages with Unsloth branded ones.
Uses a list of registered (match_text, replacement, category) rules checked in order.
"""
_rules = []
_original_showwarning = None
_installed = False
@classmethod
def add_rule(cls, match_text, replacement, category = None):
cls._rules.append((match_text, replacement, category))
if not cls._installed:
cls._install()
@classmethod
def _install(cls):
cls._original_showwarning = warnings.showwarning
cls._installed = True
def _patched_showwarning(
message, category, filename, lineno, file = None, line = None
):
msg_str = str(message)
for match_text, replacement, match_category in cls._rules:
if match_text in msg_str and (
match_category is None or category is match_category
):
print(replacement)
return
cls._original_showwarning(message, category, filename, lineno, file, line)
warnings.showwarning = _patched_showwarning
# Stop vLLM messages
if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1":
if not UNSLOTH_ENABLE_LOGGING:
try:
from vllm.worker.worker import logger as vllm_worker_logger
@ -280,6 +318,17 @@ if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1":
del vllm_executor_logger
except:
pass
try:
from vllm.v1.executor.abstract import logger as vllm_v1_executor_logger
vllm_v1_executor_logger.addFilter(HideLoggingMessage("to fall asleep"))
vllm_v1_executor_logger.addFilter(HideLoggingMessage("to wake up"))
vllm_v1_executor_logger.addFilter(
HideLoggingMessage("Executor is not sleeping")
)
del vllm_v1_executor_logger
except:
pass
try:
from vllm.core.block.prefix_caching_block import (
logger as vllm_prefix_caching_logger,
@ -519,6 +568,36 @@ class RaiseUninitialized:
transformers_logger.removeHandler(self.error_handler)
try:
from transformers.trainer import logger as transformers_trainer_logger
transformers_trainer_logger.addFilter(
HideLoggingMessage("The model is already on multiple devices.")
)
except:
pass
# Hide HF Hub unauthenticated request warnings
try:
from huggingface_hub.utils._http import logger as hf_http_logger
hf_http_logger.addFilter(
HideLoggingMessage("You are sending unauthenticated requests")
)
del hf_http_logger
except:
pass
# Replace PEFT target_parameters warning with Unsloth branded message for MoE models
ReplaceWarningMessage.add_rule(
match_text = "target_parameters",
replacement = (
"Unsloth: PEFT set target_parameters but found no matching parameters.\n"
"This is expected for MoE models - Unsloth handles MoE expert LoRA targeting separately."
),
category = RuntimeWarning,
)
# Patch get_model_param_count to record correct 4bit / 8bit
from transformers.trainer_pt_utils import is_deepspeed_zero3_enabled
@ -919,7 +998,7 @@ except ModuleNotFoundError:
xformers_attention = None
xformers_version = None
except Exception as e:
if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "0":
if UNSLOTH_ENABLE_LOGGING:
print(
"========\nSwitching to PyTorch attention since your Xformers is broken.\n========\n"
)
@ -1734,6 +1813,20 @@ def _unsloth_pre_compute_loss(self, model, inputs, *args, **kwargs):
"Using gradient accumulation will be very slightly less accurate.\n"
"Read more on gradient accumulation issues here: https://unsloth.ai/blog/gradient"
)
# Gemma3 multimodal models in transformers 5.x require token_type_ids during training.
# For text-only SFT, token_type_ids should be all zeros (no image tokens).
if "token_type_ids" not in inputs and "input_ids" in inputs:
_inner = model
for _attr in ("base_model", "model", "model"):
_inner = getattr(_inner, _attr, _inner)
if getattr(getattr(_inner, "config", None), "model_type", "") in ("gemma3",):
import sys as _sys
_mod = _sys.modules.get(type(_inner).__module__)
_has_ccm = _mod is not None and hasattr(_mod, "create_causal_mask_mapping")
if _has_ccm and _inner.training:
inputs["token_type_ids"] = torch.zeros_like(inputs["input_ids"])
outputs = self._old_compute_loss(model, inputs, *args, **kwargs)
return outputs
@ -1966,6 +2059,12 @@ def unsloth_compile_transformers(
return model_types, False
supports_sdpa = [True]
# Run patches BEFORE compiler so class replacements (e.g. GptOssTopKRouter,
# GptOssExperts) are in place before the compiler caches references to them.
for temporary_patch in TEMPORARY_PATCHES:
temporary_patch()
for model_type in model_types:
_unsloth_compile_transformers(
model_type,
@ -2608,7 +2707,7 @@ def get_moe_target_parameters(model, target_modules = None) -> Optional[List[str
if moe_params:
print(
f"Unsloth: Detected MoE model with {num_experts} experts - enabling LoRA on: {moe_params}"
f"Unsloth: Detected MoE model with {num_experts = } and {target_modules = }. Enabling LoRA on MoE parameters: {moe_params}"
)
return moe_params

View file

@ -13,6 +13,7 @@
# limitations under the License.
from .llama import *
from .llama import _get_rope_theta
from ._utils import __version__
from unsloth_zoo.utils import _get_dtype, Version
from unsloth_zoo.hf_utils import dtype_from_config
@ -256,9 +257,17 @@ class GemmaFixedRotaryEmbedding(torch.nn.Module):
config = None, # [TODO] Hack to pass in config - need to remove later
):
super().__init__()
# In transformers 5.0+, RotaryEmbedding(config) passes config as first positional arg (dim)
if (
config is None
and dim is not None
and hasattr(dim, "max_position_embeddings")
):
config = dim
dim = None
if config is not None:
# [TODO] Hack to pass in config - need to remove later
base = config.rope_theta
base = _get_rope_theta(config, default = base)
partial_rotary_factor = (
config.partial_rotary_factor
if hasattr(config, "partial_rotary_factor")
@ -433,10 +442,10 @@ class FastGemmaModel(FastLlamaModel):
return
@staticmethod
def post_patch(model, tokenizer):
def post_patch(model, tokenizer, correct_dtype = None):
# Gemma does not downcast RoPE
model, tokenizer = patch_model_and_tokenizer(
model, tokenizer, downcast_rope = False
model, tokenizer, downcast_rope = False, correct_dtype = correct_dtype
)
# Add 1 to weight

View file

@ -613,10 +613,10 @@ class FastGemma2Model(FastLlamaModel):
return
@staticmethod
def post_patch(model, tokenizer):
def post_patch(model, tokenizer, correct_dtype = None):
# Gemma does not downcast RoPE
model, tokenizer = patch_model_and_tokenizer(
model, tokenizer, downcast_rope = False
model, tokenizer, downcast_rope = False, correct_dtype = correct_dtype
)
# Add 1 to weight

View file

@ -542,7 +542,7 @@ class FastGraniteModel(FastLlamaModel):
return
@staticmethod
def post_patch(model, tokenizer):
def post_patch(model, tokenizer, correct_dtype = None):
# Torch.compile fails on embedding matrix??
# Workaround randomnly fixes it for torch versions < 2.2
model.model.embed_tokens = torch.nn.Embedding.from_pretrained(

View file

@ -26,6 +26,7 @@ from ._utils import (
_get_inference_mode_context_manager,
_prepare_model_for_qat,
)
from .loader_utils import _get_fp8_mode_and_check_settings
from ..utils.packing import (
get_packed_info_from_kwargs,
mask_packed_sequence_boundaries,
@ -207,14 +208,21 @@ def _fast_prepare_inputs_for_generation(
self,
input_ids,
attention_mask = None,
inputs_embeds = None,
**kwargs,
):
past_key_values = kwargs.get("past_key_values", None)
# 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 past_key_values is not None:
# Check for uninitialized DynamicCache
if len(past_key_values) == 0:
past_key_values = None
kwargs["past_key_values"] = None
use_inputs_embeds = inputs_embeds is not None
# New since 4.56
elif (
hasattr(past_key_values, "get_seq_length")
@ -222,9 +230,18 @@ def _fast_prepare_inputs_for_generation(
):
past_key_values = None
kwargs["past_key_values"] = None
use_inputs_embeds = inputs_embeds is not None
else:
bs, cache_length = input_ids.shape
input_ids = input_ids[:, [-1]]
if input_ids is not None and input_ids.numel() > 0:
bs, cache_length = input_ids.shape
input_ids = input_ids[:, [-1]]
device = input_ids.device
elif inputs_embeds is not None:
bs, cache_length, _ = inputs_embeds.shape
device = inputs_embeds.device
else:
bs, cache_length = 1, 0
device = "cuda" if torch.cuda.is_available() else "cpu"
# Get to the base model
base_model = self
@ -248,7 +265,7 @@ def _fast_prepare_inputs_for_generation(
"target_length": cache_length,
"dtype": self.dtype,
"cache_position": torch.arange(
cache_length, cache_length + 1, device = input_ids.device
cache_length, cache_length + 1, device = device
),
"batch_size": bs,
"config": self.config,
@ -258,7 +275,7 @@ def _fast_prepare_inputs_for_generation(
if needs_device_kw(
base_model._prepare_4d_causal_attention_mask_with_cache_position
):
kwargs["device"] = input_ids.device
kwargs["device"] = device
except:
print(
f"Unsloth: Could not inspect signature of {base_model._prepare_4d_causal_attention_mask_with_cache_position}"
@ -271,7 +288,8 @@ def _fast_prepare_inputs_for_generation(
)
)
else:
attention_mask = attention_mask[:, [-1]]
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 "
@ -282,11 +300,17 @@ def _fast_prepare_inputs_for_generation(
if "cache_position" in kwargs:
kwargs["position_ids"] = kwargs["cache_position"]
return {
"input_ids": input_ids,
result = {
"attention_mask": attention_mask,
**kwargs,
}
if use_inputs_embeds:
result["inputs_embeds"] = inputs_embeds
result["input_ids"] = None
else:
result["input_ids"] = input_ids
return result
def fix_prepare_inputs_for_generation(module):
@ -844,6 +868,11 @@ def LlamaModel_fast_forward(
input_ids = input_ids[:, : self.max_seq_length]
elif inputs_embeds is not None:
inputs_embeds = inputs_embeds[:, : self.max_seq_length, :]
if (
attention_mask is not None
and attention_mask.shape[-1] > self.max_seq_length
):
attention_mask = attention_mask[:, : self.max_seq_length]
past_key_values_length = 0
@ -1559,6 +1588,18 @@ def PeftModel_fast_forward(
)
def _get_rope_theta(config, default = 10000.0):
"""Get rope_theta from config, handling both transformers 4.x and 5.x."""
try:
return config.rope_theta
except (AttributeError, KeyError):
pass
rp = getattr(config, "rope_parameters", None)
if isinstance(rp, dict):
return rp.get("rope_theta", default)
return default
# Solves https://github.com/unslothai/unsloth/issues/168
# Static KV Cache was introduced in 4.38.0, causing training to be much slower.
# Inference can now be CUDAGraphed, but we shall retain the old rotary embeddings.
@ -1579,11 +1620,7 @@ class LlamaRotaryEmbedding(torch.nn.Module):
super().__init__()
if config is not None:
# [TODO] Hack to pass in config - need to remove later
try:
base = config.rope_theta
except:
base = getattr(config, "rope_parameters", {})
base = base["rope_theta"]
base = _get_rope_theta(config, default = base)
partial_rotary_factor = (
config.partial_rotary_factor
if hasattr(config, "partial_rotary_factor")
@ -1734,7 +1771,7 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module):
super().__init__()
if config is not None:
# [TODO] Hack to pass in config - need to remove later
base = config.rope_theta
base = _get_rope_theta(config, default = base)
partial_rotary_factor = (
config.partial_rotary_factor
if hasattr(config, "partial_rotary_factor")
@ -1870,7 +1907,7 @@ class LongRopeRotaryEmbedding(torch.nn.Module):
if config is not None:
# [TODO] Hack to pass in config - need to remove later
base = config.rope_theta
base = _get_rope_theta(config, default = base)
partial_rotary_factor = (
config.partial_rotary_factor
if hasattr(config, "partial_rotary_factor")
@ -2033,12 +2070,16 @@ def unsloth_fast_generate(
and kwargs["input_ids"] is not None
and "max_new_tokens" in kwargs
):
if (
kwargs["input_ids"].shape[-1] + kwargs["max_new_tokens"]
_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
):
raise ValueError(
f"Unsloth: input length {kwargs['input_ids'].shape[-1]} + max_new_tokens {kwargs['max_new_tokens']} exceeds the maximum sequence length of {self.config.max_position_embeddings}!\n"
f"Unsloth: input length {_ids.shape[-1]} + max_new_tokens {kwargs['max_new_tokens']} exceeds the maximum sequence length of {self.config.max_position_embeddings}!\n"
"You will need to do long context extension by increasing the `max_seq_length` in `FastLanguageModel.from_pretrained`."
)
@ -2152,6 +2193,7 @@ class FastLlamaModel:
unsloth_vllm_standby = False,
num_labels = None,
qat_scheme = None,
load_in_fp8 = False, # fp8 LoRA (True, False, 'block')
**kwargs,
):
os.environ["UNSLOTH_USE_NEW_MODEL"] = "0"
@ -2395,6 +2437,13 @@ class FastLlamaModel:
generate_batches,
)
fp8_mode = None
if load_in_fp8 != False:
fp8_mode = _get_fp8_mode_and_check_settings(
load_in_fp8,
fast_inference,
)
allowed_args = inspect.getfullargspec(load_vllm).args
load_vllm_kwargs = dict(
model_name = model_name,
@ -2408,6 +2457,7 @@ class FastLlamaModel:
disable_log_stats = disable_log_stats,
use_bitsandbytes = load_in_4bit,
unsloth_vllm_standby = unsloth_vllm_standby,
fp8_mode = fp8_mode,
)
for allowed_arg in allowed_args:
if allowed_arg not in load_vllm_kwargs and allowed_arg in kwargs:
@ -2418,7 +2468,11 @@ class FastLlamaModel:
llm = load_vllm(**load_vllm_kwargs)
# Convert to HF format
_, quant_state_dict = get_vllm_state_dict(llm, config = model_config)
_, quant_state_dict = get_vllm_state_dict(
llm,
config = model_config,
load_in_fp8 = load_in_fp8,
)
model = convert_vllm_to_huggingface(
quant_state_dict, model_config, dtype, bnb_config
)
@ -2443,7 +2497,9 @@ class FastLlamaModel:
)
model, tokenizer = patch_tokenizer(model, tokenizer)
model, tokenizer = model_patcher.post_patch(model, tokenizer)
model, tokenizer = model_patcher.post_patch(
model, tokenizer, correct_dtype = dtype
)
# Patch up QKV / O and MLP
for idx, layer in enumerate(model.model.layers):
@ -2626,9 +2682,9 @@ class FastLlamaModel:
return model, tokenizer
@staticmethod
def post_patch(model, tokenizer):
def post_patch(model, tokenizer, correct_dtype = None):
model, tokenizer = patch_model_and_tokenizer(
model, tokenizer, downcast_rope = True
model, tokenizer, downcast_rope = True, correct_dtype = correct_dtype
)
return model, tokenizer

View file

@ -37,6 +37,7 @@ from .loader_utils import (
_offline_quantize_to_fp8,
_tag_model_with_fp8_torchao_config,
get_model_name,
prepare_device_map,
)
import os, contextlib, sys
@ -99,6 +100,7 @@ global FORCE_FLOAT32
# Forces float32 precision since float16 goes to infinity
FORCE_FLOAT32 = [
"gemma3,", # Add comma bc gemma3 will match gemma3n
"gemma3text", # Gemma3TextModel (EmbeddingGemma, standalone text-only Gemma3)
"gemma3n",
"gpt_oss",
]
@ -115,6 +117,7 @@ global DISABLE_SDPA_MODEL_NAMES
# Disables some SDPA modules since it's wrong
DISABLE_SDPA_MODEL_NAMES = [
"gemma3,", # Add comma bc gemma3 will match gemma3n
"gemma3_text", # Gemma3TextModel (EmbeddingGemma) - substring match, keep underscore
]
@ -186,6 +189,16 @@ class FastLanguageModel(FastLlamaModel):
bnb_compute_dtype = getattr(torch, bnb_compute_dtype, None)
if isinstance(bnb_compute_dtype, torch.dtype):
dtype = bnb_compute_dtype
# Distributed-safe device placement for quantized models.
# In multi-GPU (torchrun), each rank must load the model on its own device
# to avoid Accelerate device relocation errors with quantized weights.
is_quantized = load_in_4bit or load_in_8bit or load_in_fp8
if is_quantized and isinstance(device_map, str):
distributed_device_map, is_dist = prepare_device_map()
if is_dist:
device_map = distributed_device_map
if load_in_8bit or full_finetuning or qat_scheme is not None:
return FastModel.from_pretrained(
model_name = model_name,
@ -277,12 +290,15 @@ class FastLanguageModel(FastLlamaModel):
load_in_4bit,
load_in_8bit,
load_in_16bit,
use_exact_model_name,
)
model_name = _offline_quantize_to_fp8(model_name, fp8_mode)
else:
assert new_model_name is not None
model_name = new_model_name
# If mapper resolved to a pre-quantized FP8 model, disable
# on-the-fly quantization to avoid double quantization
if load_in_fp8 != False and new_model_name != old_model_name:
load_in_fp8 = False
# Check if pre-quantized models are allowed
# For eg AMD Instinct GPUs need blocksize = 128, but our pre-quants are blocksize = 64
@ -602,6 +618,7 @@ class FastLanguageModel(FastLlamaModel):
random_state = random_state,
max_lora_rank = max_lora_rank,
disable_log_stats = disable_log_stats,
load_in_fp8 = load_in_fp8,
*args,
**kwargs,
)
@ -824,6 +841,16 @@ class FastModel(FastBaseModel):
)
if qat_scheme == "phone-deployment":
qat_scheme = "int8-int4"
# Distributed-safe device placement for quantized models.
# In multi-GPU (torchrun), each rank must load the model on its own device
# to avoid Accelerate device relocation errors with quantized weights.
is_quantized = load_in_4bit or load_in_8bit or load_in_fp8
if is_quantized and isinstance(device_map, str):
distributed_device_map, is_dist = prepare_device_map()
if is_dist:
device_map = distributed_device_map
# 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"):
@ -871,12 +898,15 @@ class FastModel(FastBaseModel):
load_in_4bit,
load_in_8bit,
load_in_16bit,
use_exact_model_name,
)
model_name = _offline_quantize_to_fp8(model_name, fp8_mode)
else:
assert new_model_name is not None
model_name = new_model_name
# If mapper resolved to a pre-quantized FP8 model, disable
# on-the-fly quantization to avoid double quantization
if load_in_fp8 != False and new_model_name != old_model_name:
load_in_fp8 = False
# Check if pre-quantized models are allowed
# For eg AMD Instinct GPUs need blocksize = 128, but our pre-quants are blocksize = 64
@ -1288,6 +1318,7 @@ class FastModel(FastBaseModel):
random_state = random_state,
max_lora_rank = max_lora_rank,
disable_log_stats = disable_log_stats,
load_in_fp8 = load_in_fp8,
*args,
**kwargs,
)

View file

@ -31,6 +31,7 @@ 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)
@ -117,6 +118,15 @@ def __get_model_name(
else:
if lower_model_name in FLOAT_TO_FP8_BLOCK_MAPPER:
return FLOAT_TO_FP8_BLOCK_MAPPER[lower_model_name]
# Mapper didn't find a pre-quantized model.
# For vllm >= 0.12.0, we can quantize the model to FP8 on the fly,
# so just return the original model name. Older vllm versions will
# fall through to offline quantization via _offline_quantize_to_fp8.
if importlib.util.find_spec("vllm") is not None:
import vllm
if Version(vllm.__version__) >= Version("0.12.0"):
return model_name
return None
elif not SUPPORTS_FOURBIT and lower_model_name in INT_TO_FLOAT_MAPPER:
@ -235,38 +245,12 @@ def get_model_name(model_name, load_in_4bit = True, load_in_fp8 = False):
return new_model_name if new_model_name is not None else model_name
def _get_torchao_fp8_config(fp8_mode: str):
"""
Return a `torchao.quantization.Float8DynamicActivationFloat8WeightConfig`
to be used for `load_in_fp8=True`.
"""
from torchao.quantization import (
Float8DynamicActivationFloat8WeightConfig,
PerBlock,
PerRow,
)
if fp8_mode == "row":
granularity = PerRow()
elif fp8_mode == "block":
granularity = (PerBlock([1, 128]), PerBlock([128, 128]))
else:
raise ValueError("Unsloth: `load_in_fp8` supports only 'row' or 'block'")
return Float8DynamicActivationFloat8WeightConfig(
granularity = granularity,
activation_value_lb = 1e-12,
)
def _offline_quantize_to_fp8(model_name: str, fp8_mode: str) -> str:
"""
Quantizes the model to fp8 using torchao and saving the quantized model to a
temporary location. Return the path to the quantized model.
Note: Once on-the-fly quantization is added in vllm in
https://github.com/vllm-project/vllm/pull/26327, we should
dynamically quantize the model there instead:
Note: For vllm >= 0.12.0, we should dynamically quantize the model in vllm instead:
llm = LLM(
...
@ -333,11 +317,10 @@ def _tag_model_with_fp8_torchao_config(model: torch.nn.Module, fp8_mode: str):
def _get_fp8_mode_and_check_settings(
load_in_fp8: Union[bool, str],
fast_inference: bool,
full_finetuning: bool,
load_in_4bit: bool,
load_in_8bit: bool,
load_in_16bit: bool,
use_exact_model_name: bool,
full_finetuning: bool = False,
load_in_4bit: bool = False,
load_in_8bit: bool = False,
load_in_16bit: bool = False,
) -> str:
"""
Assuming `load_in_fp8` is enabled, raise appropriate errors on incompatible settings
@ -373,8 +356,6 @@ def _get_fp8_mode_and_check_settings(
raise ValueError(
"Unsloth: `load_in_fp8` is not compatible with `load_in_4bit`, `load_in_8bit` or `load_in_16bit`",
)
if use_exact_model_name:
raise ValueError("Unsloth: `load_in_fp8` requires `use_exact_model_name=False`")
# Check if this is Hopper or above
if not (

View file

@ -70,6 +70,20 @@ except Exception:
except Exception:
trl_version = Version("0.0.0")
# Get PyTorch version for feature detection
try:
torch_version = Version(torch.__version__.split("+")[0].split("a")[0].split("b")[0])
except Exception:
torch_version = Version("0.0.0")
# Get transformers version for feature detection
try:
from transformers import __version__ as _transformers_version_raw
transformers_version = Version(_transformers_version_raw)
except Exception:
transformers_version = Version("0.0.0")
def vLLMSamplingParams(**kwargs):
from vllm import SamplingParams
@ -429,6 +443,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
for x in dir(trainer)
if x.endswith("Trainer")
and x != "Trainer"
and not x.startswith("_")
and trainer_file.split("_")[0] in x.lower()
]
config = [
@ -436,6 +451,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
for x in dir(trainer)
if x.endswith("Config")
and x != "Config"
and not x.startswith("_")
and trainer_file.split("_")[0] in x.lower()
]
if len(name) != 1:
@ -443,6 +459,47 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
f"Unsloth: Could not find Trainer class in trl.trainer.{trainer_file}. Found: {name}"
)
return
if len(config) != 1:
# TRL 0.26+: Config may be in a separate *_config.py module
config_module_name = trainer_file.replace("_trainer", "_config")
try:
config_mod = eval(f"trl.trainer.{config_module_name}")
config = [
x
for x in dir(config_mod)
if x.endswith("Config")
and x != "Config"
and not x.startswith("_")
and trainer_file.split("_")[0] in x.lower()
]
except Exception:
pass
if len(config) != 1 and len(name) == 1:
# Thin wrapper fallback: walk the Trainer's MRO to find Config
# in the real implementation module (e.g., trl.experimental.bco)
try:
_temp_cls = eval(f"trl.trainer.{trainer_file}.{name[0]}")
for _parent in _temp_cls.__mro__[1:]:
if _parent is object:
continue
_parent_mod = inspect.getmodule(_parent)
if (
_parent_mod is None
or _parent_mod.__name__ == f"trl.trainer.{trainer_file}"
):
continue
config = [
x
for x in dir(_parent_mod)
if x.endswith("Config")
and x != "Config"
and not x.startswith("_")
and trainer_file.split("_")[0] in x.lower()
]
if len(config) == 1:
break
except Exception:
pass
if len(config) != 1:
logger.info(
f"Unsloth: Could not find Config class in trl.trainer.{trainer_file}. Found: {config}"
@ -459,13 +516,38 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
f"Unsloth: Could not load {RLTrainer_name} from trl.trainer.{trainer_file}: {e}"
)
return
_config_resolved_module = None
try:
RLConfig = eval(f"trl.trainer.{trainer_file}.{RLConfig_name}")
except Exception as e:
logger.info(
f"Unsloth: Could not load {RLConfig_name} from trl.trainer.{trainer_file}: {e}"
)
return
except Exception:
# TRL 0.26+: Config may be in a separate *_config.py module
try:
config_module_name = trainer_file.replace("_trainer", "_config")
RLConfig = eval(f"trl.trainer.{config_module_name}.{RLConfig_name}")
except Exception:
# Thin wrapper fallback: load Config from parent trainer's module
_config_loaded = False
try:
_temp_cls = eval(f"trl.trainer.{trainer_file}.{name[0]}")
for _parent in _temp_cls.__mro__[1:]:
if _parent is object:
continue
_parent_mod = inspect.getmodule(_parent)
if (
_parent_mod is None
or _parent_mod.__name__ == f"trl.trainer.{trainer_file}"
):
continue
if hasattr(_parent_mod, RLConfig_name):
RLConfig = getattr(_parent_mod, RLConfig_name)
_config_resolved_module = _parent_mod
_config_loaded = True
break
except Exception:
pass
if not _config_loaded:
logger.info(f"Unsloth: Could not load {RLConfig_name}")
return
# Check name
if RLTrainer.__name__.startswith("Unsloth"):
@ -475,11 +557,66 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
print(f"Unsloth: {RLConfig.__name__} is already patched.")
return
# TRL 0.26+: Resolve thin wrappers to their experimental parent class.
# Thin wrappers are deprecation shims in trl.trainer that just forward
# *args/**kwargs to the real implementation in trl.experimental.
# Only resolve if a parent class actually lives in a trl.experimental module.
_trainer_resolved_module = None
try:
_trainer_src = inspect.getsource(RLTrainer)
_trainer_module = inspect.getmodule(RLTrainer)
_trainer_module_src = (
inspect.getsource(_trainer_module) if _trainer_module else ""
)
if (
"trl.experimental" in _trainer_src
or "trl.experimental" in _trainer_module_src
):
for _parent in RLTrainer.__mro__[1:]:
if _parent is object:
continue
_parent_mod = inspect.getmodule(_parent)
if _parent_mod is None:
continue
# Only resolve to a parent that lives in trl.experimental
if "trl.experimental" in _parent_mod.__name__:
RLTrainer = _parent
_trainer_resolved_module = _parent_mod
break
except Exception:
pass
try:
_config_src = inspect.getsource(RLConfig)
_config_module = inspect.getmodule(RLConfig)
_config_module_src = inspect.getsource(_config_module) if _config_module else ""
if (
"trl.experimental" in _config_src
or "trl.experimental" in _config_module_src
):
for _parent in RLConfig.__mro__[1:]:
if _parent is object:
continue
_parent_mod = inspect.getmodule(_parent)
if _parent_mod is None:
continue
# Only resolve to a parent that lives in trl.experimental
if "trl.experimental" in _parent_mod.__name__:
RLConfig = _parent
break
except Exception:
pass
# Get old source
old_RLTrainer_source = inspect.getsource(RLTrainer)
old_RLConfig_source = inspect.getsource(RLConfig)
all_imports = dir(trainer)
if _trainer_resolved_module is not None:
all_imports = dir(_trainer_resolved_module)
elif _config_resolved_module is not None:
all_imports = dir(_config_resolved_module)
else:
all_imports = dir(trainer)
# Fix _deprecate_arguments not getting imported so stop __ but not _
imports = [x for x in all_imports if not x.startswith("__")]
@ -830,7 +967,6 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
"per_device_train_batch_size": 4,
"gradient_accumulation_steps": 2,
"weight_decay": 0.01,
"warmup_ratio": 0.1,
"seed": 3407,
"optim": "adamw_8bit",
"learning_rate": 5e-05,
@ -857,6 +993,12 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
# "dataloader_prefetch_factor" : 2,
# "dataloader_num_workers" : 2, # Default is 0 means 1
}
# warmup_ratio deprecated in transformers >= 5.0; warmup_steps accepts float
if transformers_version >= Version("5.0.0"):
replacements["warmup_steps"] = 0.1
else:
replacements["warmup_ratio"] = 0.1
for k, v in replacements.items():
x = f"{k}( = [^,\n]{{1,}})?,\n"
y = f"'{v}'" if type(v) is str else f"{v}"
@ -926,14 +1068,15 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
# Edit dataset_num_proc
if "dataset_num_proc" in call_args:
num_proc_check = (
"if dataset_num_proc is None:\n"
"import multiprocessing as _mp\n"
"if _mp.get_start_method() != 'fork':\n"
" dataset_num_proc = None\n"
"elif dataset_num_proc is None:\n"
" import psutil\n"
" dataset_num_proc = min(max((psutil.cpu_count() or 1)+4, 2), 64)\n"
" memory_gb_left = psutil.virtual_memory().available / (1024**3)\n"
" if memory_gb_left <= 4: dataset_num_proc = 1 # Too risky, so set to 1\n"
" elif memory_gb_left <= 6: dataset_num_proc = min(2, dataset_num_proc)\n"
" elif memory_gb_left <= 10: dataset_num_proc = min(4, dataset_num_proc)\n"
" elif memory_gb_left <= 14: dataset_num_proc = min(6, dataset_num_proc)\n"
" if memory_gb_left <= 2: dataset_num_proc = 1\n"
" else: dataset_num_proc = min(dataset_num_proc, int(memory_gb_left))\n"
)
extra_args += num_proc_check
@ -1126,16 +1269,18 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
# Generate torch_compile_options based on device type
if DEVICE_TYPE == "cuda":
# CUDA-specific options (added to base options)
new_options = (
base_options
+ """
"triton.enable_persistent_tma_matmul": torch.cuda.get_device_capability()[0] >= 9,
cuda_options = """
"triton.enable_persistent_tma_matmul": torch.cuda.get_device_capability()[0] >= 9,"""
# cutlass options were added in PyTorch 2.8.0
if torch_version >= Version("2.8.0"):
cuda_options += """
"cuda.cutlass_epilogue_fusion_enabled": torch.cuda.get_device_capability()[0] >= 9,
"cuda.cutlass_tma_only": torch.cuda.get_device_capability()[0] >= 9,
"cuda.cutlass_tma_only": torch.cuda.get_device_capability()[0] >= 9,"""
cuda_options += """
"cuda.compile_opt_level" : "-O2",
"cuda.enable_cuda_lto" : True,
}"""
)
new_options = base_options + cuda_options
else:
# XPU, HIP, and other device types use base options only
new_options = (
@ -1177,18 +1322,99 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
flags = re.DOTALL,
)
# Remove TRL's unconditional bfloat16 cast of trainable params (added in
# TRL 0.26.0). TRL hardcodes bfloat16 for QLoRA per the original paper's
# recommendation, but this is wrong: it ignores the user's requested dtype
# and breaks GradScaler when training with fp16=True. Unsloth already
# handles adapter dtype correctly via patch_model_and_tokenizer, so the
# entire block is unnecessary. For GRPOTrainer the enclosing peft init
# block is already removed above, making this a no-op for GRPO.
RLTrainer_source = RLTrainer_source.replace(
'if getattr(model, "is_loaded_in_4bit", False) or getattr(model, "is_loaded_in_8bit", False):',
"if False:",
)
if RLTrainer_name == "SFTTrainer":
original_text = 'self._signature_columns = ["input_ids", "attention_mask", "completion_mask"]'
new_text = 'self._signature_columns = ["input_ids", "attention_mask", "completion_mask","labels"]'
RLTrainer_source = RLTrainer_source.replace(original_text, new_text)
# Temporary patch _is_vlm to False
# as of 0.22 it only exists in sfttrainer
original_is_vlm_text = "self._is_vlm = True"
new_is_vlm_text = "self._is_vlm = False"
RLTrainer_source = RLTrainer_source.replace(
original_is_vlm_text, new_is_vlm_text
# Do NOT override _is_vlm -- let TRL detect VLM models naturally.
# In TRL 0.27.1+, forcing _is_vlm=False causes a ValueError when
# vision datasets are used with VLM models.
#
# However, some notebooks pass a bare tokenizer (processor.tokenizer) as
# processing_class. TRL then sets _is_vlm=False even for VLM models.
# Add a model-architecture-based override before the validation check.
_vlm_check_original = (
' self._is_vision_dataset = "image" in dataset_sample or "images" in dataset_sample\n'
" if self._is_vision_dataset and not self._is_vlm:"
)
_vlm_check_patched = (
' self._is_vision_dataset = "image" in dataset_sample or "images" in dataset_sample\n'
" # Unsloth: override _is_vlm for VLM models that pass a bare tokenizer\n"
" if not self._is_vlm and self._is_vision_dataset:\n"
" _m = model\n"
' if hasattr(_m, "model"): _m = _m.model\n'
' if hasattr(getattr(_m, "config", None), "vision_config") or \\\n'
' _m.__class__.__name__.endswith("ForConditionalGeneration"):\n'
" self._is_vlm = True\n"
" if self._is_vision_dataset and not self._is_vlm:"
)
if _vlm_check_original in RLTrainer_source:
RLTrainer_source = RLTrainer_source.replace(
_vlm_check_original, _vlm_check_patched
)
# Fix TRL 0.22.x: VLM models with text-only datasets.
# TRL 0.22.x checks _is_vlm (model type) not _is_vision_dataset (dataset
# content, added in 0.25.1+). When _is_vlm=True, signature columns are
# vision-only ["messages","prompt","completion","images"], which have zero
# overlap with tokenized text columns. Fix: merge both column sets into the
# VLM branch. Extra columns not in the dataset are harmlessly ignored by
# _remove_unused_columns (it only raises when zero columns match).
_sig_vlm_old = (
'self._signature_columns = ["messages", "prompt", "completion", "images"]'
)
_sig_vlm_new = (
'self._signature_columns = ["messages", "prompt", "completion", "images",'
' "input_ids", "labels", "attention_mask", "seq_lengths", "completion_mask", "assistant_masks"]'
)
RLTrainer_source = RLTrainer_source.replace(_sig_vlm_old, _sig_vlm_new)
# Inject model reference before _prepare_dataset for dynamic
# token_type_ids detection in sft_prepare_dataset
_prep_pattern = r"([ \t]*)train_dataset = self\._prepare_dataset\("
_prep_replacement = r"\1self._unsloth_model_ref = model\n\1train_dataset = self._prepare_dataset("
RLTrainer_source = re.sub(
_prep_pattern, _prep_replacement, RLTrainer_source, count = 1
)
# Silence TRL's noisy batch_size=1 + padding-free warning (handles both
# the original "anihilate" typo and the corrected "annihilate" spelling)
for _typo in ("anihilate", "annihilate"):
_idx = RLTrainer_source.find(_typo)
if _idx == -1:
continue
# Walk backwards to find "if args.per_device_train_batch_size"
_block_start = RLTrainer_source.rfind(
"if args.per_device_train_batch_size == 1", 0, _idx
)
if _block_start == -1:
continue
# Walk backwards to the newline before the if
_line_start = RLTrainer_source.rfind("\n", 0, _block_start)
# Walk forwards past the closing paren to the end of the block
_close = RLTrainer_source.find(")", _idx)
if _close == -1:
continue
_block_end = RLTrainer_source.find("\n", _close)
if _block_end == -1:
continue
RLTrainer_source = (
RLTrainer_source[:_line_start] + RLTrainer_source[_block_end:]
)
break
# Remove multiple doc strings
if __RLConfig_doc__ != "" and RLTrainer_source.count(__RLTrainer_doc__) == 2:
@ -1198,10 +1424,16 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
RLTrainer_source = re.sub(r"[\n]{3,}", "\n", RLTrainer_source)
# Create new function
_resolved_module = _trainer_resolved_module or _config_resolved_module
_model_location = (
_resolved_module.__name__
if _resolved_module is not None
else f"trl.trainer.{trainer_file}"
)
created_module = create_new_function(
f"Unsloth{RLTrainer_name}",
RLTrainer_source,
f"trl.trainer.{trainer_file}",
_model_location,
imports,
overwrite = False,
)
@ -1454,12 +1686,15 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import
for function in functions:
if not hasattr(RLTrainer, function):
continue
fx = getattr(RLTrainer, function)
try:
source = inspect.getsource(fx)
except:
continue
original_source = source
if function in changed:
original_source, source = changed[function]
else:
fx = getattr(RLTrainer, function)
try:
source = inspect.getsource(fx)
except:
continue
original_source = source
# Check for function
for edit_function in edit_functions:
@ -1575,7 +1810,10 @@ def patch_trl_rl_trainers():
if x.islower() and x.endswith("_trainer") and x != "base_trainer"
]
for trainer in all_trainers:
_patch_trl_rl_trainers(trainer)
try:
_patch_trl_rl_trainers(trainer)
except Exception as e:
logger.warning_once(f"Unsloth: Could not patch trl.trainer.{trainer}: {e}")
return

View file

@ -217,6 +217,19 @@ def sft_trainer_compute_loss(function_name, function):
RL_FUNCTIONS["sft_trainer"].append(sft_trainer_compute_loss)
# Fix bare pop("push_to_hub_token") in compiled SFT/IterativeSFT trainer __init__
# On transformers 5.0+, to_dict() no longer includes push_to_hub_token, so bare pop KeyErrors
def sft_trainer_push_to_hub_token(function_name, function):
if function_name != "__init__":
return function
return function.replace(
'dict_args.pop("push_to_hub_token")', 'dict_args.pop("push_to_hub_token", None)'
)
RL_FUNCTIONS["sft_trainer"].append(sft_trainer_push_to_hub_token)
# Autocast precision for GRPO
def grpo_trainer__prepare_inputs(function_name, function):
if function_name != "_prepare_inputs":
@ -459,7 +472,7 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
function = function.replace(string_to_find, replacement_string)
if trl_version >= Version("0.25.0"):
if trl_version >= Version("0.24.0"):
# We replace the call using 'completions' with one using 'completions_text'
string_to_find = " rewards_per_func = self._calculate_rewards(inputs, prompts, completions, completion_ids_list)"
replacement_string = (
@ -1193,6 +1206,29 @@ def grpo_trainer_compute_loss(function_name, function):
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer_compute_loss)
# Fix KTO shape mismatch when Unsloth model forward truncates input_ids
# but labels aren't truncated. TRL 0.27.2+ _process_tokens only truncates
# completions, not prompts -- so prompts exceeding max_seq_length cause the
# model to produce shorter logits than the labels expect.
def kto_trainer_get_batch_logps(function_name, function):
if function_name != "get_batch_logps":
return function
# The raise is inside an if block inside the method, so we need
# to preserve the exact indentation of the raise statement.
old = 'raise ValueError("Logits (batch and sequence length dim) and labels must have the same shape.")'
new = (
"# Unsloth: auto-truncate to shorter sequence length (model may have truncated input_ids)\n"
" _min_len = min(logits.shape[1], labels.shape[1])\n"
" logits = logits[:, :_min_len, :]\n"
" labels = labels[:, :_min_len]"
)
function = function.replace(old, new)
return function
RL_FUNCTIONS["kto_trainer"].append(kto_trainer_get_batch_logps)
# https://github.com/huggingface/trl/blob/main/trl/trainer/grpo_trainer.py#L356
# TRL warns if batch size is not a multiple of num_generations -> fix this.
def grpo_trainer_fix_batch_size(RLTrainer_source, RLConfig_source):
@ -1267,7 +1303,7 @@ def openenv_vllm_reload_weights():
try:
import trl.experimental.openenv.utils as openenv_utils
import trl.experimental.openenv as openenv
except ImportError as e:
except (ImportError, NameError, Exception) as e:
logger.info(f"Unsloth: Failed to import trl openenv: {e}")
logger.info(
"Unsloth: trl.experimental.openenv not available — skipping RL openenv patches."

View file

@ -14,7 +14,7 @@
import logging
from .loader import FastModel
from .loader import FastModel, DISABLE_SDPA_MODEL_NAMES
from ._utils import SUPPORTS_BFLOAT16
import inspect
import json
@ -288,6 +288,237 @@ def _save_pretrained_gguf(
return result
def _push_to_hub_gguf(
self,
repo_id,
tokenizer = None,
quantization_method = "fast_quantized",
first_conversion = None,
token = None,
private = None,
commit_message = "Upload GGUF SentenceTransformer model trained with Unsloth",
commit_description = "Upload GGUF model trained with Unsloth 2x faster",
max_shard_size = "5GB",
temporary_location = "_unsloth_temporary_saved_buffers",
maximum_memory_usage = 0.85,
create_pr = False,
revision = None,
tags = None,
**kwargs,
):
"""
Converts the SentenceTransformer model to GGUF format and pushes to the Hugging Face Hub.
This method:
1. Saves the model locally to a temporary directory in GGUF format.
2. Uploads the GGUF files, config, Ollama Modelfile, and README to the Hub.
3. Cleans up the temporary directory.
Args:
repo_id (str): The Hugging Face Hub repo ID (e.g., "username/model-name").
tokenizer: The tokenizer to save. Defaults to `self.tokenizer`.
quantization_method (str or list): GGUF quantization method(s). Can be a string or list of strings.
Choose from the following options:
* "not_quantized" : Recommended. Fast conversion. Slow inference, big files.
* "fast_quantized" : Recommended. Fast conversion. OK inference, OK file size.
* "quantized" : Recommended. Slow conversion. Fast inference, small files.
* "f32" : Not recommended. Retains 100% accuracy, but super slow and memory hungry.
* "f16" : Fastest conversion + retains 100% accuracy. Slow and memory hungry.
* "q8_0" : Fast conversion. High resource use, but generally acceptable.
* "q4_k_m" : Recommended. Uses Q6_K for half of the attention.wv and feed_forward.w2 tensors, else Q4_K
* "q5_k_m" : Recommended. Uses Q6_K for half of the attention.wv and feed_forward.w2 tensors, else Q5_K
* "q2_k" : Uses Q4_K for the attention.vw and feed_forward.w2 tensors, Q2_K for the other tensors.
* "q3_k_l" : Uses Q5_K for the attention.wv, attention.wo, and feed_forward.w2 tensors, else Q3_K
* "q3_k_m" : Uses Q4_K for the attention.wv, attention.wo, and feed_forward.w2 tensors, else Q3_K
* "q3_k_s" : Uses Q3_K for all tensors
* "q4_0" : Original quant method, 4-bit.
* "q4_1" : Higher accuracy than q4_0 but not as high as q5_0. However has quicker inference than q5 models.
* "q4_k_s" : Uses Q4_K for all tensors
* "q5_0" : Higher accuracy, higher resource usage and slower inference.
* "q5_1" : Even higher accuracy, resource usage and slower inference.
* "q5_k_s" : Uses Q5_K for all tensors
* "q6_k" : Uses Q8_K for all tensors
first_conversion (str, optional): The initial conversion format before quantization.
token (str, optional): Hugging Face token. Uses cached token if not provided.
private (bool, optional): Whether the repo should be private.
commit_message (str): Commit message for the upload.
commit_description (str): Commit description for the upload.
max_shard_size (str): Maximum shard size for saving.
temporary_location (str): Temp directory for intermediate files.
maximum_memory_usage (float): Max fraction of memory to use.
create_pr (bool): Whether to create a pull request instead of pushing directly.
revision (str, optional): Branch/revision to push to.
tags (list, optional): Additional tags for the repo.
Returns:
str: The full repo ID on Hugging Face Hub.
"""
if token is None:
token = get_token()
if token is None:
raise ValueError(
"No HF token provided. Please provide a token or login with `huggingface-cli login`"
)
api = HfApi(token = token)
# Determine full repo_id
if "/" not in repo_id:
username = api.whoami()["name"]
full_repo_id = f"{username}/{repo_id}"
else:
full_repo_id = repo_id
model_name = full_repo_id.split("/")[-1]
# Create repo
try:
api.create_repo(
repo_id = full_repo_id,
private = private,
exist_ok = True,
repo_type = "model",
)
except Exception as e:
print(f"Unsloth Warning: Could not create repo: {e}")
# Save to temporary directory first
with tempfile.TemporaryDirectory(prefix = "unsloth_st_gguf_") as temp_dir:
print(f"Unsloth: Converting SentenceTransformer to GGUF format...")
# Call save_pretrained_gguf to do the local conversion
result = _save_pretrained_gguf(
self,
save_directory = temp_dir,
tokenizer = tokenizer,
quantization_method = quantization_method,
first_conversion = first_conversion,
push_to_hub = False, # We handle upload ourselves
token = token,
max_shard_size = max_shard_size,
temporary_location = temporary_location,
maximum_memory_usage = maximum_memory_usage,
)
gguf_files = result.get("gguf_files", [])
modelfile_location = result.get("modelfile_location", None)
is_vlm = result.get("is_vlm", False)
fix_bos_token = result.get("fix_bos_token", False)
print(f"Unsloth: Uploading GGUF to https://huggingface.co/{full_repo_id}...")
# Upload GGUF files
for file_location in gguf_files:
if os.path.exists(file_location):
filename = os.path.basename(file_location)
print(f" Uploading {filename}...")
api.upload_file(
path_or_fileobj = file_location,
path_in_repo = filename,
repo_id = full_repo_id,
repo_type = "model",
commit_message = commit_message,
commit_description = commit_description,
create_pr = create_pr,
revision = revision,
)
# Upload Modelfile if exists
if modelfile_location and os.path.exists(modelfile_location):
print(" Uploading Ollama Modelfile...")
api.upload_file(
path_or_fileobj = modelfile_location,
path_in_repo = "Modelfile",
repo_id = full_repo_id,
repo_type = "model",
commit_message = f"{commit_message} - Ollama Modelfile",
create_pr = create_pr,
revision = revision,
)
# Upload config.json if exists
config_path = os.path.join(temp_dir, "config.json")
if os.path.exists(config_path):
print(" Uploading config.json...")
api.upload_file(
path_or_fileobj = config_path,
path_in_repo = "config.json",
repo_id = full_repo_id,
repo_type = "model",
commit_message = f"{commit_message} - config",
create_pr = create_pr,
revision = revision,
)
# Create and upload README
gguf_basenames = [os.path.basename(f) for f in gguf_files if os.path.exists(f)]
readme_content = f"""---
tags:
- gguf
- llama.cpp
- unsloth
- sentence-transformers
{"- vision-language-model" if is_vlm else ""}
---
# {model_name} - GGUF
This sentence-transformers model was finetuned and converted to GGUF format using [Unsloth](https://github.com/unslothai/unsloth).
## Available Model files:
"""
for fname in gguf_basenames:
readme_content += f"- `{fname}`\n"
if modelfile_location and os.path.exists(modelfile_location):
readme_content += "\n## Ollama\n"
readme_content += "An Ollama Modelfile is included for easy deployment.\n"
if fix_bos_token:
readme_content += "\n## Note\n"
readme_content += (
"The model's BOS token behavior was adjusted for GGUF compatibility.\n"
)
readme_content += (
"\nThis was trained 2x faster with [Unsloth](https://github.com/unslothai/unsloth)\n"
'[<img src="https://raw.githubusercontent.com/unslothai/unsloth/main/images/unsloth%20made%20with%20love.png" width="200"/>](https://github.com/unslothai/unsloth)\n'
)
readme_path = os.path.join(temp_dir, "README.md")
with open(readme_path, "w", encoding = "utf-8") as f:
f.write(readme_content)
api.upload_file(
path_or_fileobj = readme_path,
path_in_repo = "README.md",
repo_id = full_repo_id,
repo_type = "model",
commit_message = "Add README",
create_pr = create_pr,
revision = revision,
)
# Add tags
all_tags = ["gguf", "llama-cpp", "unsloth", "sentence-transformers"]
if is_vlm:
all_tags.append("vision-language-model")
if tags is not None:
if isinstance(tags, (list, tuple)):
all_tags.extend(tags)
else:
all_tags.append(tags)
try:
api.add_tags(repo_id = full_repo_id, tags = all_tags, repo_type = "model")
except:
pass
print(
f"Unsloth: Successfully uploaded GGUF to https://huggingface.co/{full_repo_id}"
)
return full_repo_id
class FastSentenceTransformer(FastModel):
@staticmethod
def _read_pooling_mode(model_name, token):
@ -1230,8 +1461,17 @@ class FastSentenceTransformer(FastModel):
model_kwargs = {"torch_dtype": dtype}
# Enable SDPA if supported (1.2x extra speedup on top of torch.compile)
# But disable for models with known SDPA + torch.compile backward issues
_force_eager = False
for _sdpa_model in DISABLE_SDPA_MODEL_NAMES:
if _sdpa_model in model_type.lower():
supports_sdpa = False
_force_eager = True
break
if supports_sdpa:
model_kwargs["attn_implementation"] = "sdpa"
elif _force_eager:
model_kwargs["attn_implementation"] = "eager"
# Print optimization status
sdpa_str = " + SDPA" if supports_sdpa else ""
@ -1321,6 +1561,8 @@ class FastSentenceTransformer(FastModel):
_save_pretrained_gguf, st_model
)
st_model.push_to_hub_gguf = types.MethodType(_push_to_hub_gguf, st_model)
def _push_to_hub_merged(self, repo_id, **push_kwargs):
hub_token = push_kwargs.get("token", None) or get_token()
if hub_token is None:
@ -1523,6 +1765,8 @@ class FastSentenceTransformer(FastModel):
_save_pretrained_gguf, st_model
)
st_model.push_to_hub_gguf = types.MethodType(_push_to_hub_gguf, st_model)
def _push_to_hub_merged(self, repo_id, **kwargs):
token = kwargs.get("token", None) or get_token()
if token is None:
@ -1845,6 +2089,20 @@ def _patch_sentence_transformer_trainer():
# Call original __init__
_original_init(self, *args, **kwargs)
# Disable mixed precision when FORCE_FLOAT32 is active (matches rl.py behavior)
if os.environ.get("UNSLOTH_FORCE_FLOAT32", "0") == "1":
if hasattr(self, "args") and self.args is not None:
if self.args.fp16 or self.args.bf16:
print(
"Unsloth: Switching to float32 training since model cannot work with float16"
)
self.args.fp16 = False
self.args.bf16 = False
if hasattr(self.args, "bf16_full_eval"):
self.args.bf16_full_eval = False
if hasattr(self.args, "fp16_full_eval"):
self.args.fp16_full_eval = False
SentenceTransformerTrainer.__init__ = _patched_init
SentenceTransformerTrainer._unsloth_auto_compile_patched = True

View file

@ -31,6 +31,7 @@ from ..kernels import (
)
from ._utils import __version__, importlib_version, _prepare_model_for_qat
from ._utils import *
from .loader_utils import _get_fp8_mode_and_check_settings
from ..save import patch_saving_functions
from ..models.loader_utils import is_distributed
from unsloth_zoo.gradient_checkpointing import (
@ -317,6 +318,91 @@ def unsloth_base_fast_generate(
return output
def _construct_vlm_processor_fallback(
tokenizer_name, model_type, token, trust_remote_code
):
"""Construct a VLM processor manually when AutoProcessor.from_pretrained fails.
Some VLMs (e.g., LFM2.5-VL) have tokenizer_class entries that AutoTokenizer
cannot resolve. This function loads the image processor and tokenizer separately,
sets required special token attributes, and constructs the processor.
"""
try:
from transformers import AutoImageProcessor, PreTrainedTokenizerFast, AutoConfig
from transformers.models.auto.processing_auto import PROCESSOR_MAPPING_NAMES
import json
# Load image processor
image_processor = AutoImageProcessor.from_pretrained(
tokenizer_name,
token = token,
trust_remote_code = trust_remote_code,
)
# Load tokenizer via PreTrainedTokenizerFast (bypasses tokenizer_class check)
tok = PreTrainedTokenizerFast.from_pretrained(
tokenizer_name,
padding_side = "left",
token = token,
trust_remote_code = trust_remote_code,
)
# Read tokenizer_config.json for model-specific special tokens
try:
from huggingface_hub import hf_hub_download
config_path = hf_hub_download(
tokenizer_name, "tokenizer_config.json", token = token
)
with open(config_path, "r", encoding = "utf-8") as f:
tok_config = json.load(f)
# Set model-specific special tokens and their IDs
for key in (
"image_token",
"image_start_token",
"image_end_token",
"image_thumbnail",
"video_token",
):
if key in tok_config and not hasattr(tok, key):
setattr(tok, key, tok_config[key])
id_key = key + "_id" if not key.endswith("_id") else key
token_id = tok.convert_tokens_to_ids(tok_config[key])
if not hasattr(tok, id_key):
setattr(tok, id_key, token_id)
except Exception:
pass
# Find the processor class - try model_type first, then top-level config model_type
proc_class_name = PROCESSOR_MAPPING_NAMES.get(model_type)
if proc_class_name is None:
# model_type might be a sub-model type (e.g. "lfm2" instead of "lfm2_vl").
# Try the top-level config.model_type which often has the processor mapping.
try:
config = AutoConfig.from_pretrained(
tokenizer_name,
token = token,
trust_remote_code = trust_remote_code,
)
proc_class_name = PROCESSOR_MAPPING_NAMES.get(config.model_type)
except Exception:
pass
if proc_class_name is not None:
import transformers
proc_class = getattr(transformers, proc_class_name, None)
if proc_class is not None:
processor = proc_class(image_processor = image_processor, tokenizer = tok)
# Copy chat_template from tokenizer to processor if needed
if not getattr(processor, "chat_template", None) and getattr(
tok, "chat_template", None
):
processor.chat_template = tok.chat_template
return processor
except Exception:
pass
return None
class FastBaseModel:
@staticmethod
def from_pretrained(
@ -348,6 +434,7 @@ class FastBaseModel:
max_lora_rank = 64,
disable_log_stats = False,
unsloth_vllm_standby = False,
load_in_fp8 = False, # fp8 LoRA (True, False, 'block')
**kwargs,
):
if unsloth_vllm_standby and os.environ.get("UNSLOTH_VLLM_STANDBY", "0") != "1":
@ -753,6 +840,17 @@ class FastBaseModel:
model_name, model_config
)
fp8_mode = None
if load_in_fp8 != False:
fp8_mode = _get_fp8_mode_and_check_settings(
load_in_fp8,
fast_inference,
full_finetuning,
load_in_4bit,
load_in_8bit,
load_in_16bit,
)
allowed_args = inspect.getfullargspec(load_vllm).args
load_vllm_kwargs = dict(
model_name = model_name,
@ -767,6 +865,7 @@ class FastBaseModel:
use_bitsandbytes = load_in_4bit,
unsloth_vllm_standby = unsloth_vllm_standby,
is_vision_model = is_vlm,
fp8_mode = fp8_mode,
)
for allowed_arg in allowed_args:
if allowed_arg not in load_vllm_kwargs and allowed_arg in kwargs:
@ -780,6 +879,7 @@ class FastBaseModel:
llm,
config = model_config,
is_vision_model = is_vlm,
load_in_fp8 = load_in_fp8,
)
model = convert_vllm_to_huggingface(
quant_state_dict,
@ -826,14 +926,17 @@ class FastBaseModel:
if (whisper_language and whisper_task) or auto_model.__name__.endswith(
"ForConditionalGeneration"
):
tokenizer = auto_processor.from_pretrained(
tokenizer_name,
padding_side = "left",
token = token,
language = whisper_language,
task = whisper_task,
trust_remote_code = trust_remote_code,
)
try:
tokenizer = auto_processor.from_pretrained(
tokenizer_name,
padding_side = "left",
token = token,
language = whisper_language,
task = whisper_task,
trust_remote_code = trust_remote_code,
)
except Exception:
tokenizer = None
else:
try:
tokenizer = auto_processor.from_pretrained(
@ -849,6 +952,23 @@ class FastBaseModel:
token = token,
trust_remote_code = trust_remote_code,
)
# If processor loading failed (e.g., tokenizer class not found),
# try constructing the processor manually from separate components.
if tokenizer is None and is_vlm:
tokenizer = _construct_vlm_processor_fallback(
tokenizer_name,
model_type_arch,
token,
trust_remote_code,
)
if tokenizer is None:
import sys
print(
f"Unsloth: Warning - VLM processor fallback returned None for model_type={model_type_arch}",
file = sys.stderr,
)
if hasattr(tokenizer, "tokenizer"):
__tokenizer = tokenizer.tokenizer
# Add padding side as well
@ -872,7 +992,29 @@ class FastBaseModel:
do_forced_float32 = do_forced_float32,
correct_dtype = correct_dtype,
)
model, tokenizer = patch_tokenizer(model, tokenizer)
try:
model, tokenizer = patch_tokenizer(model, tokenizer)
except Exception as _patch_err:
# Some VLM processors (e.g., ERNIE VL) may fail during tokenizer patching.
# Try loading tokenizer separately via AutoTokenizer as fallback.
try:
from transformers import AutoTokenizer as _AutoTokenizer
_fallback_tok = _AutoTokenizer.from_pretrained(
tokenizer_name,
padding_side = "left",
token = token,
trust_remote_code = trust_remote_code,
)
model, _fallback_tok = patch_tokenizer(model, _fallback_tok)
# Re-attach as processor wrapper if original was a processor
if hasattr(tokenizer, "image_processor"):
tokenizer.tokenizer = _fallback_tok
else:
tokenizer = _fallback_tok
except Exception:
# If fallback also fails, raise the original error
raise _patch_err
model = post_patch_loss_function(model)
# Log Unsloth version for future fastpaths for inference
@ -880,10 +1022,31 @@ class FastBaseModel:
model.config.update({"unsloth_version": __version__})
patch_saving_functions(model, vision = True)
if tokenizer is None:
del model
raise RuntimeError(
"Unsloth: The tokenizer is weirdly not loaded? Please check if there is one."
)
# Last resort: try loading tokenizer via AutoTokenizer, then PreTrainedTokenizerFast
try:
from transformers import AutoTokenizer as _AutoTokenizer
tokenizer = _AutoTokenizer.from_pretrained(
tokenizer_name,
padding_side = "left",
token = token,
trust_remote_code = trust_remote_code,
)
except Exception:
try:
from transformers import PreTrainedTokenizerFast
tokenizer = PreTrainedTokenizerFast.from_pretrained(
tokenizer_name,
padding_side = "left",
token = token,
trust_remote_code = trust_remote_code,
)
except Exception:
del model
raise RuntimeError(
"Unsloth: The tokenizer is weirdly not loaded? Please check if there is one."
)
patch_saving_functions(tokenizer, vision = True)
# Fix gradient accumulation

View file

@ -806,7 +806,7 @@ TEMPLATE """<|system|>
PARAMETER stop "<|system|>"
PARAMETER stop "<|user|>"
PARAMETER stop "<|assistant|>"
PARAMETER "</s>"
PARAMETER stop "</s>"
SYSTEM """You are a helpful AI assistant."""
'''
@ -974,6 +974,7 @@ TEMPLATE """{{- range $i, $_ := .Messages }}
OLLAMA_TEMPLATES["qwen-25"] = qwen25_ollama
OLLAMA_TEMPLATES["qwen-2.5"] = qwen25_ollama
OLLAMA_TEMPLATES["qwen-25-coder"] = qwen_25_coder_ollama
OLLAMA_TEMPLATES["qwen-25-vl"] = qwen_25_vl_ollama
OLLAMA_TEMPLATES["openthinker"] = openthinker_ollama

View file

@ -1257,6 +1257,16 @@ def save_to_gguf(
"Please check disk space and try again."
)
# Move initial GGUF files into a dedicated _gguf directory
gguf_directory = f"{model_directory}_gguf"
os.makedirs(gguf_directory, exist_ok = True)
moved_files = []
for fpath in initial_files:
dst = os.path.join(gguf_directory, os.path.basename(fpath))
shutil.move(fpath, dst)
moved_files.append(dst)
initial_files = moved_files
print(f"Unsloth: Initial conversion completed! Files: {initial_files}")
# Step 4: Additional quantizations using llama-quantize
@ -1276,8 +1286,9 @@ def save_to_gguf(
print(
f"Unsloth: [2] Converting GGUF {first_conversion_dtype} into {quant_method}. This might take 10 minutes..."
)
output_location = f"{model_name}.{quant_method.upper()}.gguf"
output_location = os.path.join(
gguf_directory, f"{model_name}.{quant_method.upper()}.gguf"
)
try:
# Use the quantize_gguf function we created
quantized_file = quantize_gguf(
@ -1316,7 +1327,7 @@ def save_to_gguf(
print("Unsloth: Model files cleanup...")
if quants_created:
all_saved_locations.remove(base_gguf)
Path(base_gguf).unlink()
Path(base_gguf).unlink(missing_ok = True)
# flip the list to get [text_model, mmproj] order. for text models stays the same.
all_saved_locations.reverse()
@ -1996,6 +2007,7 @@ def unsloth_save_pretrained_gguf(
raise RuntimeError(f"Unsloth: GGUF conversion failed: {e}")
# Step 9: Create Ollama modelfile
gguf_directory = f"{save_directory}_gguf"
modelfile_location = None
ollama_success = False
if all_file_locations:
@ -2004,13 +2016,12 @@ def unsloth_save_pretrained_gguf(
modelfile = create_ollama_modelfile(tokenizer, base_model_name, ".")
else:
modelfile = create_ollama_modelfile(
tokenizer, base_model_name, all_file_locations[0]
tokenizer,
base_model_name,
os.path.basename(all_file_locations[0]),
)
if modelfile is not None:
if is_vlm_update:
modelfile_location = os.path.join(save_directory, "Modelfile")
else:
modelfile_location = os.path.join(os.getcwd(), "Modelfile")
modelfile_location = os.path.join(gguf_directory, "Modelfile")
with open(modelfile_location, "w", encoding = "utf-8") as file:
file.write(modelfile)
ollama_success = True
@ -2027,28 +2038,25 @@ def unsloth_save_pretrained_gguf(
if is_vlm_update:
print("\n")
print(
f"Unsloth: example usage for Multimodal LLMs: llama-mtmd-cli -m {all_file_locations[0]} --mmproj {all_file_locations[-1]}"
f"Unsloth: example usage for Multimodal LLMs: llama.cpp/llama-mtmd-cli -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-cli --model {all_file_locations[0]} -p "why is the sky blue?"'
f'Unsloth: example usage for text only LLMs: llama.cpp/llama-cli --model {all_file_locations[0]} -p "why is the sky blue?"'
)
if ollama_success and is_vlm_update:
if ollama_success:
print(f"Unsloth: Saved Ollama Modelfile to {modelfile_location}")
print(
"Unsloth: convert model to ollama format by running - ollama create model_name -f ./Modelfile - inside save directory."
)
if ollama_success and not is_vlm_update:
print("Unsloth: Saved Ollama Modelfile to current directory")
print(
"Unsloth: convert model to ollama format by running - ollama create model_name -f ./Modelfile - inside current directory."
f"Unsloth: convert model to ollama format by running - ollama create model_name -f {modelfile_location}"
)
# Return a dict with all needed info for push_to_hub
return {
"save_directory": save_directory,
"gguf_directory": gguf_directory,
"gguf_files": all_file_locations,
"modelfile_location": modelfile_location,
"want_full_precision": want_full_precision,
@ -2148,10 +2156,11 @@ def unsloth_push_to_hub_gguf(
if cleanup_temp:
import shutil
try:
shutil.rmtree(save_directory)
except:
pass
for d in [save_directory, f"{save_directory}_gguf"]:
try:
shutil.rmtree(d)
except:
pass
raise RuntimeError(f"Failed to convert model to GGUF: {e}")
# Step 3: Upload to HuggingFace Hub
@ -2334,14 +2343,16 @@ This model was finetuned and converted to GGUF format using [Unsloth](https://gi
finally:
# Clean up temporary directory
if cleanup_temp and os.path.exists(save_directory):
if cleanup_temp:
print("Unsloth: Cleaning up temporary files...")
import shutil
try:
shutil.rmtree(save_directory)
except:
pass
for d in [save_directory, f"{save_directory}_gguf"]:
if os.path.exists(d):
try:
shutil.rmtree(d)
except:
pass
return full_repo_id

View file

@ -1021,7 +1021,10 @@ def patch_sft_trainer_tokenizer():
"kto_trainer.KTOTrainer",
):
function_name, replacer = "train", "if resume_from_checkpoint is False:"
function = getsource(eval(f"trl.trainer.{path_to_trainer}.{function_name}"))
try:
function = getsource(eval(f"trl.trainer.{path_to_trainer}.{function_name}"))
except Exception:
continue
where = function.find("def")
function = function.split("\n")
function = "\n".join(x[where:] for x in function)

View file

@ -200,13 +200,57 @@ class UnslothTrainer(SFTTrainer):
# From `trl>=0.13.0`, they changed how to pass several params to the trainer
# We need to patch to make the transition smooth
def _resolve_trainer_params(trainer_class, init_fn):
"""Resolve the real named parameters for a trainer __init__.
Some TRL trainers (e.g., ORPOTrainer in TRL 0.27.1) are thin wrappers
with only ``def __init__(self, *args, **kwargs)``. For those, walk the
MRO and return the first parent class that has real named parameters.
"""
params = inspect.signature(init_fn).parameters
named = {
k
for k, v in params.items()
if v.kind
in (inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.KEYWORD_ONLY)
and k != "self"
}
if named:
return set(params.keys())
# Thin wrapper detected - walk MRO for real signature
for cls in trainer_class.__mro__[1:]:
if cls is object:
continue
parent_init = cls.__dict__.get("__init__")
if parent_init is None:
continue
try:
parent_params = inspect.signature(parent_init).parameters
parent_named = {
k
for k, v in parent_params.items()
if v.kind
in (
inspect.Parameter.POSITIONAL_OR_KEYWORD,
inspect.Parameter.KEYWORD_ONLY,
)
and k != "self"
}
if parent_named:
return set(parent_params.keys())
except (ValueError, TypeError):
continue
return set(params.keys())
def _backwards_compatible_trainer(trainer_class, config_class):
original_init = trainer_class.__init__
@wraps(original_init)
def new_init(self, *args, **kwargs):
# All Trainer tokenizer are now called processing_class
trainer_params = set(inspect.signature(original_init).parameters.keys())
trainer_params = _resolve_trainer_params(trainer_class, original_init)
if "processing_class" in trainer_params and "tokenizer" in kwargs:
kwargs["processing_class"] = kwargs.pop("tokenizer")

View file

@ -107,12 +107,14 @@ def configure_sample_packing(config):
_ensure_trl_warning_filter()
setattr(config, "packing", True)
setattr(config, "padding_free", True)
setattr(config, "remove_unused_columns", False)
def configure_padding_free(config):
"""Mutate an ``SFTConfig`` so TRL enables padding-free batching without packing."""
_ensure_trl_warning_filter()
setattr(config, "padding_free", True)
setattr(config, "remove_unused_columns", False)
def enable_sample_packing(
@ -151,6 +153,12 @@ def enable_sample_packing(
lengths = example.get(sequence_lengths_key)
if isinstance(lengths, Iterable):
seq_lengths.extend(int(length) for length in lengths)
# Fallback: infer lengths from tokenized inputs when metadata is absent
if not seq_lengths:
for example in examples:
ids = example.get("input_ids")
if isinstance(ids, Iterable):
seq_lengths.append(len(ids))
if seq_lengths:
batch["packed_seq_lengths"] = torch.tensor(
seq_lengths, dtype = torch.int32
@ -176,6 +184,8 @@ def enable_padding_free_metadata(model, trainer):
mark_allow_overlength(model)
if hasattr(collator, "return_position_ids"):
collator.return_position_ids = True
if hasattr(trainer, "args") and hasattr(trainer.args, "remove_unused_columns"):
trainer.args.remove_unused_columns = False
original_torch_call = collator.torch_call