Compare commits
2 commits
main
...
dh/recover
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
76054a562a | ||
|
|
2d7d2b7dea |
6 changed files with 233 additions and 39 deletions
|
|
@ -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]",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
178
unsloth/models/mlx_model.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue