diff --git a/README.md b/README.md index ca5b6533b5..4d934eb033 100644 --- a/README.md +++ b/README.md @@ -180,7 +180,8 @@ python -m bitsandbytes - We're in 🤗Hugging Face's official docs! Check out the [SFT docs](https://huggingface.co/docs/trl/main/en/sft_trainer#accelerate-fine-tuning-2x-using-unsloth) and [DPO docs](https://huggingface.co/docs/trl/main/en/dpo_trainer#accelerate-dpo-fine-tuning-using-unsloth)! ```python -from unsloth import FastLanguageModel +from unsloth import FastLanguageModel +from unsloth import is_bfloat16_supported import torch from trl import SFTTrainer from transformers import TrainingArguments @@ -238,8 +239,8 @@ trainer = SFTTrainer( gradient_accumulation_steps = 4, warmup_steps = 10, max_steps = 60, - fp16 = not torch.cuda.is_bf16_supported(), - bf16 = torch.cuda.is_bf16_supported(), + fp16 = not is_bfloat16_supported(), + bf16 = is_bfloat16_supported(), logging_steps = 1, output_dir = "outputs", optim = "adamw_8bit", @@ -263,6 +264,7 @@ We're in 🤗Hugging Face's official docs! We're on the [SFT docs](https://huggi ```python from unsloth import FastLanguageModel, PatchDPOTrainer +from unsloth import is_bfloat16_supported PatchDPOTrainer() import torch from transformers import TrainingArguments @@ -298,8 +300,8 @@ dpo_trainer = DPOTrainer( gradient_accumulation_steps = 8, warmup_ratio = 0.1, num_train_epochs = 3, - fp16 = not torch.cuda.is_bf16_supported(), - bf16 = torch.cuda.is_bf16_supported(), + fp16 = not is_bfloat16_supported(), + bf16 = is_bfloat16_supported(), logging_steps = 1, optim = "adamw_8bit", seed = 42, diff --git a/unsloth/__init__.py b/unsloth/__init__.py index d4ca45d7d1..2dcf1e6a43 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -114,3 +114,4 @@ from .models import * from .save import * from .chat_templates import * from .tokenizer_utils import * +from .trainer import * diff --git a/unsloth/models/__init__.py b/unsloth/models/__init__.py index ff7129e06a..e67a9e5fad 100644 --- a/unsloth/models/__init__.py +++ b/unsloth/models/__init__.py @@ -17,3 +17,4 @@ from .llama import FastLlamaModel from .mistral import FastMistralModel from .qwen2 import FastQwen2Model from .dpo import PatchDPOTrainer +from ._utils import is_bfloat16_supported diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index a53de42c09..2c1eb4d5c4 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -35,7 +35,10 @@ __version__ = "2024.5" # Get Flash Attention v2 if Ampere (RTX 30xx, A100) major_version, minor_version = torch.cuda.get_device_capability() +SUPPORTS_BFLOAT16 = False + if major_version >= 8: + SUPPORTS_BFLOAT16 = True try: from flash_attn import flash_attn_func # Check for CUDA linking errors "undefined symbol: _ZNK3c106SymIntltEl" @@ -72,6 +75,10 @@ __all__ = [ "patch_tokenizer", "get_statistics", "Unsloth_Offloaded_Gradient_Checkpointer", + "offload_to_disk", + "offload_input_embeddings", + "offload_output_embeddings", + "is_bfloat16_supported", ] @@ -421,3 +428,52 @@ except: "Luckily, your training run will still work in the meantime!" ) pass + + +# Offloading to disk for modules (lm_head, embed_tokens) +import os +import pickle + +def offload_to_disk(W, model, name, temporary_location : str = "_unsloth_temporary_saved_buffers"): + file_location = os.path.join(temporary_location, model.config._name_or_path) + if not os.path.exists(file_location): + os.makedirs(file_location) + pass + + filename = os.path.join(file_location, f"{name}.pt") + W = W.weight if hasattr(W, "weight") else W + torch.save(W, filename, pickle_module = pickle, pickle_protocol = pickle.HIGHEST_PROTOCOL,) + offloaded_W = torch.load(filename, map_location = "cpu", mmap = True) + offloaded_W._offloaded_file_location = filename + return offloaded_W +pass + + +def offload_input_embeddings(model, temporary_location : str = "_unsloth_temporary_saved_buffers"): + offloaded_W = offload_to_disk(model.get_input_embeddings(), model, "input_embeddings", temporary_location) + new_input_embeddings = torch.nn.Embedding.from_pretrained(offloaded_W) + new_input_embeddings._offloaded_file_location = offloaded_W._offloaded_file_location + model.set_input_embeddings(new_input_embeddings) + return +pass + + +def offload_output_embeddings(model, temporary_location : str = "_unsloth_temporary_saved_buffers"): + offloaded_W = offload_to_disk(model.get_output_embeddings(), model, "output_embeddings", temporary_location) + + new_output_embeddings = torch.nn.Linear(1, 1, bias = None) + del new_output_embeddings.weight + new_output_embeddings.weight = offloaded_W + new_output_embeddings.in_features = offloaded_W.shape[1] + new_output_embeddings.out_features = offloaded_W.shape[0] + + new_output_embeddings._offloaded_file_location = offloaded_W._offloaded_file_location + model.set_output_embeddings(new_output_embeddings) + return +pass + + +# Fixes a weird Torch 2.3 bug which says T4s have bfloat16 +def is_bfloat16_supported(): + return SUPPORTS_BFLOAT16 +pass diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index f7fd5f13f4..1d6a282a55 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -13,6 +13,7 @@ # limitations under the License. import torch +import gc from typing import Optional, Tuple, List, Union from torch.nn.functional import scaled_dot_product_attention from transformers.models.llama.modeling_llama import ( @@ -1046,7 +1047,7 @@ class FastLlamaModel: token = os.environ["HUGGINGFACE_TOKEN"] if model_patcher is None: model_patcher = FastLlamaModel - SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported() + SUPPORTS_BFLOAT16 = is_bfloat16_supported() gpu_stats = torch.cuda.get_device_properties(0) max_memory = round(gpu_stats.total_memory / 1024 / 1024 / 1024, 3) @@ -1193,7 +1194,11 @@ class FastLlamaModel: f"O^O/ \\_/ \\ Batch size per device = {self._train_batch_size:,} | Gradient Accumulation steps = {args.gradient_accumulation_steps}\\n"\\ f"\\ / Total batch size = {total_train_batch_size:,} | Total steps = {max_steps:,}\\n"\\ f' "-____-" Number of trainable parameters = {get_model_param_count(model, trainable_only=True):,}' - logger.warning_once(debug_info)""" + logger.warning_once(debug_info) + import gc + for _ in range(3): + gc.collect() + torch.cuda.empty_cache()""" debug_info = debug_info.split('\n') debug_info = "\n".join([debug_info[0]] + [spaces + x[8:] for x in debug_info[1:]]) @@ -1370,7 +1375,6 @@ class FastLlamaModel: pass # Clear deleted GPU items - import gc for _ in range(3): gc.collect() torch.cuda.empty_cache() @@ -1396,6 +1400,7 @@ class FastLlamaModel: modules_to_save = None, init_lora_weights = True, loftq_config = {}, + temporary_location = "_unsloth_temporary_saved_buffers", **kwargs, ): transformers_set_seed(random_state) @@ -1490,19 +1495,19 @@ class FastLlamaModel: final_modules = [] for module in target_modules: if module == "lm_head": - logger.warning_once( - "Unsloth: `lm_head` should be placed in `modules_to_save` and not `target_modules`. "\ - "Luckily, we shall do it for you!" - ) + # logger.warning_once( + # "Unsloth: `lm_head` should be placed in `modules_to_save` and not `target_modules`. "\ + # "Luckily, we shall do it for you!" + # ) train_lm_head = True if modules_to_save is None: modules_to_save = ["lm_head"] else: modules_to_save.append("lm_head") elif module == "embed_tokens": - logger.warning_once( - "Unsloth: `embed_tokens` should be placed in `modules_to_save` and not `target_modules`. "\ - "Luckily, we shall do it for you!" - ) + # logger.warning_once( + # "Unsloth: `embed_tokens` should be placed in `modules_to_save` and not `target_modules`. "\ + # "Luckily, we shall do it for you!" + # ) train_embed_tokens = True if modules_to_save is None: modules_to_save = ["embed_tokens"] else: modules_to_save.append("embed_tokens") @@ -1579,6 +1584,35 @@ class FastLlamaModel: _saved_temp_tokenizer = model._saved_temp_tokenizer lora_config = LoraConfig(**arguments) + + # First offload lm_head and embed_tokens to disk + input_embeddings_device = model. get_input_embeddings().weight.device + output_embeddings_device = model.get_output_embeddings().weight.device + + if use_gradient_checkpointing == "unsloth": + if train_embed_tokens: + print("Unsloth: Offloading input_embeddings to disk to save VRAM") + offload_input_embeddings(model, temporary_location) + pass + + # Remove old items to save VRAM + for _ in range(3): + gc.collect() + torch.cuda.empty_cache() + pass + + if train_lm_head: + print("Unsloth: Offloading output_embeddings to disk to save VRAM") + offload_output_embeddings(model, temporary_location) + pass + + # Remove old items to save VRAM + for _ in range(3): + gc.collect() + torch.cuda.empty_cache() + pass + pass + model = _get_peft_model(model, lora_config) model._saved_temp_tokenizer = _saved_temp_tokenizer @@ -1589,14 +1623,16 @@ class FastLlamaModel: if train_embed_tokens: print("Unsloth: Casting embed_tokens to float32") assert(hasattr(model.model.model.embed_tokens, "modules_to_save")) - model.model.model.embed_tokens.modules_to_save.default.to(torch.float32) + model.model.model.embed_tokens.modules_to_save.default\ + .to(device = input_embeddings_device, dtype = torch.float32, non_blocking = True) model.model.model.embed_tokens.modules_to_save.default.requires_grad_(True) pass if train_lm_head: print("Unsloth: Casting lm_head to float32") assert(hasattr(model.model.lm_head, "modules_to_save")) - model.model.lm_head.modules_to_save.default.to(torch.float32) + model.model.lm_head.modules_to_save.default\ + .to(device = output_embeddings_device, dtype = torch.float32, non_blocking = True) model.model.lm_head.modules_to_save.default.requires_grad_(True) pass @@ -1612,6 +1648,12 @@ class FastLlamaModel: internal_model._saved_temp_tokenizer.padding_side = "right" pass + # Clear deleted GPU items + for _ in range(3): + gc.collect() + torch.cuda.empty_cache() + pass + return model pass @@ -1715,7 +1757,7 @@ class FastLlamaModel: n_mlp += 1 else: logger.warning_once( - "Unsloth cannot patch MLP layers with our manual autograd engine since either LoRA adapters\n"\ + "Not an error, but Unsloth cannot patch MLP layers with our manual autograd engine since either LoRA adapters\n"\ "are not enabled or a bias term (like in Qwen) is used." ) pass @@ -1738,7 +1780,7 @@ class FastLlamaModel: n_qkv += 1 else: logger.warning_once( - "Unsloth cannot patch Attention layers with our manual autograd engine since either LoRA adapters\n"\ + "Not an error, but Unsloth cannot patch Attention layers with our manual autograd engine since either LoRA adapters\n"\ "are not enabled or a bias term (like in Qwen) is used." ) pass @@ -1753,7 +1795,7 @@ class FastLlamaModel: n_o += 1 else: logger.warning_once( - "Unsloth cannot patch O projection layer with our manual autograd engine since either LoRA adapters\n"\ + "Not an error, but Unsloth cannot patch O projection layer with our manual autograd engine since either LoRA adapters\n"\ "are not enabled or a bias term (like in Qwen) is used." ) pass diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index 4594919b38..365d60a3e4 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -314,7 +314,7 @@ class FastMistralModel(FastLlamaModel): logger.warning_once("Unsloth: Mistral models do not support RoPE scaling.") pass - SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported() + SUPPORTS_BFLOAT16 = is_bfloat16_supported() gpu_stats = torch.cuda.get_device_properties(0) max_memory = round(gpu_stats.total_memory / 1024 / 1024 / 1024, 3) diff --git a/unsloth/trainer.py b/unsloth/trainer.py new file mode 100644 index 0000000000..b234a98d8b --- /dev/null +++ b/unsloth/trainer.py @@ -0,0 +1,95 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from dataclasses import dataclass, field +from typing import Optional +from transformers import TrainingArguments +from trl import SFTTrainer + +__all__ = [ + "UnslothTrainingArguments", + "UnslothTrainer", +] + + +@dataclass +class UnslothTrainingArguments(TrainingArguments): + embedding_learning_rate : Optional[float] = field( + default = None, + metadata = {"help" : "Different learning rates for embeddings and lm_head."} + ) +pass + + +def _create_unsloth_optimizer( + model, + optimizer_cls, + optimizer_kwargs, + embedding_lr = 5e-5, +): + lr = optimizer_kwargs["lr"] + weight_decay = optimizer_kwargs.get("weight_decay", 0.0) + + param_groups = \ + { + "non_embeddings" : {}, + "embeddings" : {}, + } + + for name, param in model.named_parameters(): + if not param.requires_grad: continue + if name.endswith("modules_to_save.default.weight"): + partial_name = name[:-len(".modules_to_save.default.weight")] + partial_name = partial_name[partial_name.rfind(".")+1:] + print(f"Unsloth: Setting lr = {embedding_lr:.2e} instead of {lr:.2e} for {partial_name}.") + param_groups["embeddings"] [name] = param + else: + param_groups["non_embeddings"][name] = param + pass + pass + + optimizer_grouped_parameters = [ + { + "params" : list(param_groups["non_embeddings"].values()), + "weight_decay" : weight_decay, + "lr" : lr, + }, + { + "params" : list(param_groups["embeddings"].values()), + "weight_decay" : weight_decay, + "lr" : embedding_lr, + }, + ] + optimizer = optimizer_cls(optimizer_grouped_parameters, **optimizer_kwargs) + return optimizer +pass + + +class UnslothTrainer(SFTTrainer): + def create_optimizer(self): + embedding_learning_rate = getattr(self.args, "embedding_learning_rate", None) + if embedding_learning_rate is None: return super().create_optimizer() + + if self.optimizer is None: + optimizer_cls, optimizer_kwargs = SFTTrainer.get_optimizer_cls_and_kwargs(self.args) + self.optimizer = _create_unsloth_optimizer( + self.model, + optimizer_cls, + optimizer_kwargs, + embedding_learning_rate, + ) + pass + return self.optimizer + pass +pass