Nightly (#506)
* Update llama.py * offload * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * continued pretraining trainer * Update trainer.py * Update trainer.py * Update trainer.py * Update trainer.py * is_bfloat16_supported * Update __init__.py * Update README.md * Update llama.py
This commit is contained in:
parent
0eb57384bd
commit
5fc27f1f96
7 changed files with 219 additions and 22 deletions
12
README.md
12
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,
|
||||
|
|
|
|||
|
|
@ -114,3 +114,4 @@ from .models import *
|
|||
from .save import *
|
||||
from .chat_templates import *
|
||||
from .tokenizer_utils import *
|
||||
from .trainer import *
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
95
unsloth/trainer.py
Normal file
95
unsloth/trainer.py
Normal file
|
|
@ -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
|
||||
Loading…
Add table
Add a link
Reference in a new issue