Compare commits

...
Sign in to create a new pull request.

2 commits

Author SHA1 Message Date
pre-commit-ci[bot]
76054a562a [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-03-12 23:10:51 +00:00
JINO-ROHIT
2d7d2b7dea feat: add mlx model and trainer 2026-03-12 23:10:51 +00:00
6 changed files with 233 additions and 39 deletions

View file

@ -82,6 +82,11 @@ windows = [
"bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0 ; (sys_platform == 'win32')",
"xformers>=0.0.22.post7 ; (sys_platform == 'win32')",
]
mac = [
"unsloth[huggingface]",
"mlx>=0.12.0 ; sys_platform == 'darwin' and platform_machine == 'arm64'",
"mlx-lm>=0.9.0 ; sys_platform == 'darwin' and platform_machine == 'arm64'",
]
base = [
"unsloth[huggingface]",
]

View file

@ -223,9 +223,9 @@ elif DEVICE_TYPE == "xpu":
# For Gradio HF Spaces?
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
import triton
if DEVICE_TYPE == "cuda":
import triton
libcuda_dirs = lambda: None
if Version(triton.__version__) >= Version("3.0.0"):
try:
@ -308,23 +308,22 @@ elif DEVICE_TYPE == "xpu":
# TODO: check triton for intel installed properly.
pass
from .models import *
from .models import __version__
from .save import *
from .chat_templates import *
from .tokenizer_utils import *
from .trainer import *
elif DEVICE_TYPE != "mps":
from .models import *
from .models import __version__
from .save import *
from .chat_templates import *
from .tokenizer_utils import *
from .trainer import *
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
from unsloth_zoo.rl_environments import (
check_python_modules,
create_locked_down_function,
execute_with_time_limit,
Benchmarker,
is_port_open,
launch_openenv,
)
# Export dataprep utilities for CLI and downstream users
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
from unsloth_zoo.rl_environments import (
check_python_modules,
create_locked_down_function,
execute_with_time_limit,
Benchmarker,
is_port_open,
launch_openenv,
)
# Patch TRL trainers for backwards compatibility
_patch_trl_trainer()
# Patch TRL trainers for backwards compatibility
_patch_trl_trainer()

View file

@ -41,6 +41,8 @@ def get_device_type():
return "cuda"
elif hasattr(torch, "xpu") and torch.xpu.is_available():
return "xpu"
elif hasattr(torch, "mps") and torch.mps.is_available():
return "mps"
# Check torch.accelerator
if hasattr(torch, "accelerator"):
if not torch.accelerator.is_available():
@ -48,14 +50,14 @@ def get_device_type():
"Unsloth cannot find any torch accelerator? You need a GPU."
)
accelerator = str(torch.accelerator.current_accelerator())
if accelerator in ("cuda", "xpu", "hip"):
if accelerator in ("cuda", "xpu", "hip", "mps"):
raise RuntimeError(
f"Unsloth: Weirdly `torch.cuda.is_available()`, `torch.xpu.is_available()` and `is_hip` all failed.\n"
f"But `torch.accelerator.current_accelerator()` works with it being = `{accelerator}`\n"
f"Please reinstall torch - it's most likely broken :("
)
raise NotImplementedError(
"Unsloth currently only works on NVIDIA, AMD and Intel GPUs."
"Unsloth currently only works on NVIDIA, AMD, Intel GPUs, MAC Silicon and MLX."
)
@ -64,6 +66,8 @@ DEVICE_TYPE: str = get_device_type()
DEVICE_TYPE_TORCH = DEVICE_TYPE
if DEVICE_TYPE_TORCH == "hip":
DEVICE_TYPE_TORCH = "cuda"
elif DEVICE_TYPE_TORCH == "mps":
DEVICE_TYPE_TORCH = "mps"
@functools.cache

View file

@ -12,20 +12,25 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from .llama import FastLlamaModel
from .loader import FastLanguageModel, FastVisionModel, FastTextModel, FastModel
from .mistral import FastMistralModel
from .qwen2 import FastQwen2Model
from .qwen3 import FastQwen3Model
from .qwen3_moe import FastQwen3MoeModel
from .granite import FastGraniteModel
from .sentence_transformer import FastSentenceTransformer
from ..device_type import DEVICE_TYPE
try:
from .falcon_h1 import FastFalconH1Model
except:
# transformers_version < 4.53.0 does not have falcon_h1 so silently skip it for now
pass
from .dpo import PatchDPOTrainer, PatchKTOTrainer
from ._utils import is_bfloat16_supported, is_vLLM_available, __version__
from .rl import PatchFastRL, vLLMSamplingParams
if DEVICE_TYPE != "mps":
from .llama import FastLlamaModel
from .loader import FastLanguageModel, FastVisionModel, FastTextModel, FastModel
from .mistral import FastMistralModel
from .qwen2 import FastQwen2Model
from .qwen3 import FastQwen3Model
from .qwen3_moe import FastQwen3MoeModel
from .granite import FastGraniteModel
from .sentence_transformer import FastSentenceTransformer
try:
from .falcon_h1 import FastFalconH1Model
except:
# transformers_version < 4.53.0 does not have falcon_h1 so silently skip it for now
pass
from .dpo import PatchDPOTrainer, PatchKTOTrainer
from ._utils import is_bfloat16_supported, is_vLLM_available, __version__
from .rl import PatchFastRL, vLLMSamplingParams
else:
from .mlx_model import FastMLXModel

View file

@ -812,6 +812,9 @@ elif DEVICE_TYPE == "xpu":
else:
torch_amp_custom_fwd = torch.amp.custom_fwd(device_type = "xpu")
torch_amp_custom_bwd = torch.amp.custom_bwd(device_type = "xpu")
else:
torch_amp_custom_fwd = None
torch_amp_custom_bwd = None
# =============================================
# =============================================

178
unsloth/models/mlx_model.py Normal file
View file

@ -0,0 +1,178 @@
import os
import json
from typing import Optional, Dict, Any, Union, Tuple
from dataclasses import dataclass
from mlx_lm import load
from mlx_lm.tuner import (
train,
TrainingArgs,
datasets,
linear_to_lora_layers,
)
import mlx.optimizers as optim
from mlx.utils import tree_flatten
from ..device_type import DEVICE_TYPE
@dataclass
class MLXTrainingArguments:
"""training arguments for MLX models."""
adapter_file: str = "adapters.safetensors"
max_seq_length: int = 2048
grad_checkpoint: bool = True
grad_accumulation_steps: int = 1
iters: int = 100
batch_size: int = 4
val_batches: int = 10
def to_dict(self) -> Dict[str, Any]:
return {
"adapter_file": self.adapter_file,
"max_seq_length": self.max_seq_length,
"grad_checkpoint": self.grad_checkpoint,
"grad_accumulation_steps": self.grad_accumulation_steps,
"iters": self.iters,
"batch_size": self.batch_size,
"val_batches": self.val_batches,
}
class MLXLoraConfig:
def __init__(
self,
rank: int = 8,
scale: float = 20.0,
dropout: float = 0.0,
num_layers: int = 8,
):
self.rank = rank
self.scale = scale
self.dropout = dropout
self.num_layers = num_layers
def to_dict(self) -> Dict[str, Any]:
return {
"num_layers": self.num_layers,
"lora_parameters": {
"rank": self.rank,
"scale": self.scale,
"dropout": self.dropout,
},
}
def save(self, adapter_path: str):
os.makedirs(adapter_path, exist_ok = True)
config_path = os.path.join(adapter_path, "adapter_config.json")
with open(config_path, "w") as f:
json.dump(self.to_dict(), f, indent = 4)
class MLXTrainer:
def prepare_model_for_training(
self,
model: Any,
lora_config: Optional[MLXLoraConfig] = None,
) -> Any:
if lora_config is None:
lora_config = MLXLoraConfig()
model.freeze()
linear_to_lora_layers(
model,
lora_config.num_layers,
lora_config.to_dict()["lora_parameters"],
)
num_train_params = sum(
v.size for _, v in tree_flatten(model.trainable_parameters())
)
print(f"number of trainable parameters: {num_train_params}")
model.train()
return model
def _train(
self,
model: Any,
training_args: Union[MLXTrainingArguments, Dict[str, Any]],
train_dataset: Any,
val_dataset: Any = None,
learning_rate: float = 1e-5,
):
if isinstance(training_args, MLXTrainingArguments):
args_dict = training_args.to_dict()
else:
args_dict = training_args
args = TrainingArgs(**args_dict)
optimizer = optim.Adam(learning_rate = learning_rate)
train_set = datasets.CacheDataset(train_dataset)
val_set = datasets.CacheDataset(val_dataset) if val_dataset else None
train(
model = model,
args = args,
optimizer = optimizer,
train_dataset = train_set,
val_dataset = val_set,
)
class FastMLXModel:
@staticmethod
def from_pretrained(
model_name: str,
**kwargs,
) -> Tuple[Any, Any]:
print(f"Unsloth: Loading model with MLX: {model_name}")
model, tokenizer = load(model_name)
return model, tokenizer
@staticmethod
def for_inference(
model_name: str,
adapter_path: Optional[str] = None,
) -> Any:
if adapter_path:
model, _ = load(model_name, adapter_path = adapter_path)
else:
model, _ = load(model_name)
return model
@staticmethod
def train(
model: Any,
train_set: Any,
val_set: Any,
lora_config: Optional[MLXLoraConfig] = None,
iterations: int = 100,
learning_rate: float = 1e-5,
):
if DEVICE_TYPE != "mps":
raise RuntimeError("This function requires running on Apple Silicon")
trainer = MLXTrainer()
if lora_config is None:
lora_config = MLXLoraConfig()
trainer.prepare_model_for_training(model, lora_config)
trainer._train(
model = model,
training_args = MLXTrainingArguments(iters = iterations),
train_dataset = train_set,
val_dataset = val_set,
learning_rate = learning_rate,
)
return model