Bug fixes (#1891)
* Update rl.py * Patching * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * NEFTune * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Extra replacements * Update rl_replacements.py * Update rl.py * extra RL replacements * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update _utils.py * Update loader_utils.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * autocast * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update pyproject.toml * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update _utils.py * Update llama.py * Update _utils.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * GRPO optimized * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Selective Log softmax * Fix GRPO bsz * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Fix TRL * Metrics GRPO * Update rl_replacements.py * Update rl_replacements.py * No compile * Update rl.py * Remove docs * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * llama-quantize on WINDOWS WSL error fix - edit save.py (gguf saving breaks) (#1649) * edit save.py to fix gguf saving breaks. * add check for .exe or not exe file extension for linux and windows * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * unsloth_num_chunks * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py (#1754) Fix typo in comment: know -> now. This was printed when running the Llama3.1_(8B)-GRPO.ipynb example notebook, so I'd expect others to run into it as well. * Optional logits * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * fix an import error (#1767) * fix an import error * Delete .gitignore * Update loader.py * Update save.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * SamplingParams * Convert mask to float (#1762) * [Windows Support] Add latest `xformers` wheels to pyproject.toml (#1753) * Add latest xformers * Add a couple of lines to docs * vLLMSamplingParams * Update __init__.py * default num_chunks == -1 * Versioning * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update _utils.py * Update rl_replacements.py * Update rl_replacements.py * Update pyproject.toml * Update pyproject.toml * Export Model to ollama.com (#1648) * Ollama Export Model to ollama.com Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> * Check for model_name Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> * subprocess use instead of requests | added check for ollama server Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> * create_ollama_model Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> * create_ollama_model | fix Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> * Push to Ollama Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> --------- Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> * Update cross_entropy_loss.py * torch_cuda_device * Update utils.py * Update utils.py * Update utils.py * device * device * Update loader.py * Update llama.py * Update README.md * Update llama.py * Update llama.py * Update _utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * __version__ * Update rl.py * Bug fixes --------- Signed-off-by: Jyotin Goel <b22ai063@iitj.ac.in> Co-authored-by: Gennadii Manzhos <105049664+everythingisc00l@users.noreply.github.com> Co-authored-by: Seth Weidman <seth@sethweidman.com> Co-authored-by: Nino Risteski <95188570+NinoRisteski@users.noreply.github.com> Co-authored-by: Edd <68678137+Erland366@users.noreply.github.com> Co-authored-by: Ben <6579034+versipellis@users.noreply.github.com> Co-authored-by: Jyotin Goel <120490013+gjyotin305@users.noreply.github.com>
This commit is contained in:
parent
c3f80d7cd7
commit
7cef895be7
16 changed files with 400 additions and 246 deletions
62
README.md
62
README.md
|
|
@ -232,10 +232,8 @@ print(f'pip install --upgrade pip && pip install "unsloth[{x}] @ git+https://git
|
|||
|
||||
```python
|
||||
from unsloth import FastLanguageModel
|
||||
from unsloth import is_bfloat16_supported
|
||||
import torch
|
||||
from trl import SFTTrainer
|
||||
from transformers import TrainingArguments
|
||||
from trl import SFTTrainer, SFTConfig
|
||||
from datasets import load_dataset
|
||||
max_seq_length = 2048 # Supports RoPE Scaling interally, so choose any!
|
||||
# Get LAION dataset
|
||||
|
|
@ -244,21 +242,28 @@ dataset = load_dataset("json", data_files = {"train" : url}, split = "train")
|
|||
|
||||
# 4bit pre quantized models we support for 4x faster downloading + no OOMs.
|
||||
fourbit_models = [
|
||||
"unsloth/mistral-7b-v0.3-bnb-4bit", # New Mistral v3 2x faster!
|
||||
"unsloth/Meta-Llama-3.1-8B-bnb-4bit", # Llama-3.1 2x faster
|
||||
"unsloth/Meta-Llama-3.1-8B-Instruct-bnb-4bit",
|
||||
"unsloth/Meta-Llama-3.1-70B-bnb-4bit",
|
||||
"unsloth/Meta-Llama-3.1-405B-bnb-4bit", # 4bit for 405b!
|
||||
"unsloth/Mistral-Small-Instruct-2409", # Mistral 22b 2x faster!
|
||||
"unsloth/mistral-7b-instruct-v0.3-bnb-4bit",
|
||||
"unsloth/llama-3-8b-bnb-4bit", # Llama-3 15 trillion tokens model 2x faster!
|
||||
"unsloth/llama-3-8b-Instruct-bnb-4bit",
|
||||
"unsloth/llama-3-70b-bnb-4bit",
|
||||
"unsloth/Phi-3-mini-4k-instruct", # Phi-3 2x faster!
|
||||
"unsloth/Phi-3.5-mini-instruct", # Phi-3.5 2x faster!
|
||||
"unsloth/Phi-3-medium-4k-instruct",
|
||||
"unsloth/mistral-7b-bnb-4bit",
|
||||
"unsloth/gemma-7b-bnb-4bit", # Gemma 2.2x faster!
|
||||
"unsloth/gemma-2-9b-bnb-4bit",
|
||||
"unsloth/gemma-2-27b-bnb-4bit", # Gemma 2x faster!
|
||||
|
||||
"unsloth/Llama-3.2-1B-bnb-4bit", # NEW! Llama 3.2 models
|
||||
"unsloth/Llama-3.2-1B-Instruct-bnb-4bit",
|
||||
"unsloth/Llama-3.2-3B-bnb-4bit",
|
||||
"unsloth/Llama-3.2-3B-Instruct-bnb-4bit",
|
||||
|
||||
"unsloth/Llama-3.3-70B-Instruct-bnb-4bit" # NEW! Llama 3.3 70B!
|
||||
] # More models at https://huggingface.co/unsloth
|
||||
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
model_name = "unsloth/llama-3-8b-bnb-4bit",
|
||||
model_name = "unsloth/Llama-3.2-1B",
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = None,
|
||||
load_in_4bit = True,
|
||||
)
|
||||
|
||||
|
|
@ -282,16 +287,14 @@ model = FastLanguageModel.get_peft_model(
|
|||
trainer = SFTTrainer(
|
||||
model = model,
|
||||
train_dataset = dataset,
|
||||
dataset_text_field = "text",
|
||||
max_seq_length = max_seq_length,
|
||||
tokenizer = tokenizer,
|
||||
args = TrainingArguments(
|
||||
args = SFTConfig(
|
||||
dataset_text_field = "text",
|
||||
max_seq_length = max_seq_length,
|
||||
per_device_train_batch_size = 2,
|
||||
gradient_accumulation_steps = 4,
|
||||
warmup_steps = 10,
|
||||
max_steps = 60,
|
||||
fp16 = not is_bfloat16_supported(),
|
||||
bf16 = is_bfloat16_supported(),
|
||||
logging_steps = 1,
|
||||
output_dir = "outputs",
|
||||
optim = "adamw_8bit",
|
||||
|
|
@ -323,17 +326,14 @@ RL including DPO, GRPO, PPO, Reward Modelling, Online DPO all work with Unsloth.
|
|||
import os
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = "0" # Optional set GPU device ID
|
||||
|
||||
from unsloth import FastLanguageModel, PatchDPOTrainer
|
||||
from unsloth import is_bfloat16_supported
|
||||
PatchDPOTrainer()
|
||||
from unsloth import FastLanguageModel
|
||||
import torch
|
||||
from transformers import TrainingArguments
|
||||
from trl import DPOTrainer
|
||||
from trl import DPOTrainer, DPOConfig
|
||||
max_seq_length = 2048
|
||||
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
model_name = "unsloth/zephyr-sft-bnb-4bit",
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = None,
|
||||
load_in_4bit = True,
|
||||
)
|
||||
|
||||
|
|
@ -355,24 +355,22 @@ model = FastLanguageModel.get_peft_model(
|
|||
dpo_trainer = DPOTrainer(
|
||||
model = model,
|
||||
ref_model = None,
|
||||
args = TrainingArguments(
|
||||
train_dataset = YOUR_DATASET_HERE,
|
||||
# eval_dataset = YOUR_DATASET_HERE,
|
||||
tokenizer = tokenizer,
|
||||
args = DPOConfig(
|
||||
per_device_train_batch_size = 4,
|
||||
gradient_accumulation_steps = 8,
|
||||
warmup_ratio = 0.1,
|
||||
num_train_epochs = 3,
|
||||
fp16 = not is_bfloat16_supported(),
|
||||
bf16 = is_bfloat16_supported(),
|
||||
logging_steps = 1,
|
||||
optim = "adamw_8bit",
|
||||
seed = 42,
|
||||
output_dir = "outputs",
|
||||
max_length = 1024,
|
||||
max_prompt_length = 512,
|
||||
beta = 0.1,
|
||||
),
|
||||
beta = 0.1,
|
||||
train_dataset = YOUR_DATASET_HERE,
|
||||
# eval_dataset = YOUR_DATASET_HERE,
|
||||
tokenizer = tokenizer,
|
||||
max_length = 1024,
|
||||
max_prompt_length = 512,
|
||||
)
|
||||
dpo_trainer.train()
|
||||
```
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ triton = [
|
|||
]
|
||||
|
||||
windows=[
|
||||
"unsloth_zoo>=2025.2.7",
|
||||
"unsloth_zoo>=2025.3.1",
|
||||
"packaging",
|
||||
"tyro",
|
||||
"transformers>=4.46.1,!=4.47.0",
|
||||
|
|
@ -61,7 +61,7 @@ windows=[
|
|||
"xformers>=0.0.22.post7 ; platform_system == 'Windows'",
|
||||
]
|
||||
huggingface = [
|
||||
"unsloth_zoo>=2025.2.7",
|
||||
"unsloth_zoo>=2025.3.1",
|
||||
"packaging",
|
||||
"tyro",
|
||||
"transformers>=4.46.1,!=4.47.0",
|
||||
|
|
|
|||
|
|
@ -198,7 +198,7 @@ pass
|
|||
# Check for unsloth_zoo
|
||||
try:
|
||||
unsloth_zoo_version = importlib_version("unsloth_zoo")
|
||||
if Version(unsloth_zoo_version) < Version("2025.2.6"):
|
||||
if Version(unsloth_zoo_version) < Version("2025.3.1"):
|
||||
try:
|
||||
os.system("pip install --upgrade --no-cache-dir --no-deps unsloth_zoo")
|
||||
except:
|
||||
|
|
@ -212,6 +212,7 @@ except:
|
|||
pass
|
||||
|
||||
from .models import *
|
||||
from .models import __version__
|
||||
from .save import *
|
||||
from .chat_templates import *
|
||||
from .tokenizer_utils import *
|
||||
|
|
|
|||
|
|
@ -15,7 +15,13 @@
|
|||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
from .utils import calculate_settings, MAX_FUSED_SIZE, triton_tanh, triton_cast
|
||||
from .utils import (
|
||||
calculate_settings,
|
||||
MAX_FUSED_SIZE,
|
||||
triton_tanh,
|
||||
triton_cast,
|
||||
torch_cuda_device,
|
||||
)
|
||||
from transformers.models.llama.modeling_llama import logger
|
||||
from packaging.version import Version
|
||||
|
||||
|
|
@ -279,10 +285,11 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
n_rows : int
|
||||
vocab_size : int
|
||||
n_rows, vocab_size = logits.shape
|
||||
device = logits.device
|
||||
|
||||
div, mod = divmod(vocab_size, MAX_FUSED_SIZE)
|
||||
n_chunks : int = div + (mod != 0)
|
||||
losses = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0")
|
||||
losses = torch.empty(n_rows, dtype = torch.float32, device = device)
|
||||
|
||||
DO_SOFTCAPPING : bool = bool(logit_softcapping != 0)
|
||||
DO_LOGIT_SCALING : bool = bool(logit_scaling != 0)
|
||||
|
|
@ -292,39 +299,41 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
if n_chunks == 1:
|
||||
# For small vocabs <= 65336 like Llama, Mistral
|
||||
BLOCK_SIZE, num_warps = calculate_settings(vocab_size)
|
||||
logsumexp = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0")
|
||||
logsumexp = torch.empty(n_rows, dtype = torch.float32, device = device)
|
||||
|
||||
_cross_entropy_forward[(n_rows,)](
|
||||
logits, logits.stride(0),
|
||||
losses,
|
||||
logsumexp,
|
||||
labels,
|
||||
VOCAB_SIZE = vocab_size,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
DO_SOFTCAPPING = DO_SOFTCAPPING,
|
||||
SOFTCAP = logit_softcapping,
|
||||
DO_LOGIT_SCALING = DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = logit_scaling,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
with torch_cuda_device(device):
|
||||
_cross_entropy_forward[(n_rows,)](
|
||||
logits, logits.stride(0),
|
||||
losses,
|
||||
logsumexp,
|
||||
labels,
|
||||
VOCAB_SIZE = vocab_size,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
DO_SOFTCAPPING = DO_SOFTCAPPING,
|
||||
SOFTCAP = logit_softcapping,
|
||||
DO_LOGIT_SCALING = DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = logit_scaling,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
else:
|
||||
# For large vocabs > 65336 like Gemma 256K
|
||||
logsumexp = torch.empty((n_rows, n_chunks,), dtype = torch.float32, device = "cuda:0")
|
||||
logsumexp = torch.empty((n_rows, n_chunks,), dtype = torch.float32, device = device)
|
||||
|
||||
_chunked_cross_entropy_forward[(n_rows, n_chunks,)](
|
||||
logits, logits.stride(0),
|
||||
losses,
|
||||
logsumexp,
|
||||
labels,
|
||||
VOCAB_SIZE = vocab_size,
|
||||
N_CHUNKS = n_chunks,
|
||||
BLOCK_SIZE = MAX_FUSED_SIZE,
|
||||
DO_SOFTCAPPING = DO_SOFTCAPPING,
|
||||
SOFTCAP = logit_softcapping,
|
||||
DO_LOGIT_SCALING = DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = logit_scaling,
|
||||
num_warps = 32,
|
||||
)
|
||||
with torch_cuda_device(device):
|
||||
_chunked_cross_entropy_forward[(n_rows, n_chunks,)](
|
||||
logits, logits.stride(0),
|
||||
losses,
|
||||
logsumexp,
|
||||
labels,
|
||||
VOCAB_SIZE = vocab_size,
|
||||
N_CHUNKS = n_chunks,
|
||||
BLOCK_SIZE = MAX_FUSED_SIZE,
|
||||
DO_SOFTCAPPING = DO_SOFTCAPPING,
|
||||
SOFTCAP = logit_softcapping,
|
||||
DO_LOGIT_SCALING = DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = logit_scaling,
|
||||
num_warps = 32,
|
||||
)
|
||||
# logsumexp(chunked_logsumexp) - x
|
||||
# Do the -x separately
|
||||
logsumexp = torch.logsumexp(logsumexp, dim = 1) # Row sum
|
||||
|
|
@ -354,19 +363,20 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
div, mod = divmod(vocab_size, BLOCK_SIZE)
|
||||
n_blocks : int = div + (mod != 0)
|
||||
|
||||
_cross_entropy_backward[(n_rows, n_blocks,)](
|
||||
logits, logits.stride(0),
|
||||
dlosses, dlosses.stride(0),
|
||||
logsumexp,
|
||||
labels,
|
||||
VOCAB_SIZE = vocab_size,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
DO_SOFTCAPPING = ctx.DO_SOFTCAPPING,
|
||||
SOFTCAP = ctx.logit_softcapping,
|
||||
DO_LOGIT_SCALING = ctx.DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = ctx.logit_scaling,
|
||||
num_warps = 8,
|
||||
)
|
||||
with torch_cuda_device(dlosses.device):
|
||||
_cross_entropy_backward[(n_rows, n_blocks,)](
|
||||
logits, logits.stride(0),
|
||||
dlosses, dlosses.stride(0),
|
||||
logsumexp,
|
||||
labels,
|
||||
VOCAB_SIZE = vocab_size,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
DO_SOFTCAPPING = ctx.DO_SOFTCAPPING,
|
||||
SOFTCAP = ctx.logit_softcapping,
|
||||
DO_LOGIT_SCALING = ctx.DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = ctx.logit_scaling,
|
||||
num_warps = 8,
|
||||
)
|
||||
return logits, None, None, None,
|
||||
pass
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -15,7 +15,11 @@
|
|||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
from .utils import calculate_settings, triton_tanh
|
||||
from .utils import (
|
||||
calculate_settings,
|
||||
triton_tanh,
|
||||
torch_cuda_device,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
|
|
@ -41,9 +45,11 @@ pass
|
|||
def geglu_exact_forward_kernel(gate, up):
|
||||
batch, seq_len, hd = gate.shape
|
||||
n_elements = gate.numel()
|
||||
out = torch.empty((batch, seq_len, hd), dtype = gate.dtype, device = "cuda:0")
|
||||
device = gate.device
|
||||
out = torch.empty((batch, seq_len, hd), dtype = gate.dtype, device = device)
|
||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||
_exact_forward_kernel[grid](gate, up, out, n_elements, BLOCK_SIZE = 1024,)
|
||||
with torch_cuda_device(device):
|
||||
_exact_forward_kernel[grid](gate, up, out, n_elements, BLOCK_SIZE = 1024,)
|
||||
return out
|
||||
pass
|
||||
|
||||
|
|
@ -99,7 +105,8 @@ def geglu_exact_backward_kernel(DW, e, g):
|
|||
batch_seq_len, hd = e.shape
|
||||
n_elements = e.numel()
|
||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||
_exact_backward_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
||||
with torch_cuda_device(e.device):
|
||||
_exact_backward_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
||||
return DW, e, g
|
||||
pass
|
||||
|
||||
|
|
@ -133,9 +140,11 @@ pass
|
|||
def geglu_approx_forward_kernel(gate, up):
|
||||
batch, seq_len, hd = gate.shape
|
||||
n_elements = gate.numel()
|
||||
out = torch.empty((batch, seq_len, hd), dtype = gate.dtype, device = "cuda:0")
|
||||
device = gate.device
|
||||
out = torch.empty((batch, seq_len, hd), dtype = gate.dtype, device = device)
|
||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||
_approx_forward_kernel[grid](gate, up, out, n_elements, BLOCK_SIZE = 1024,)
|
||||
with torch_cuda_device(device):
|
||||
_approx_forward_kernel[grid](gate, up, out, n_elements, BLOCK_SIZE = 1024,)
|
||||
return out
|
||||
pass
|
||||
|
||||
|
|
@ -198,6 +207,7 @@ def geglu_approx_backward_kernel(DW, e, g):
|
|||
batch_seq_len, hd = e.shape
|
||||
n_elements = e.numel()
|
||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||
_approx_backward_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
||||
with torch_cuda_device(e.device):
|
||||
_approx_backward_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
||||
return DW, e, g
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@
|
|||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
from .utils import calculate_settings
|
||||
from .utils import calculate_settings, torch_cuda_device
|
||||
from unsloth_zoo.patching_utils import (
|
||||
patch_layernorm,
|
||||
)
|
||||
|
|
@ -111,17 +111,18 @@ class Fast_Layernorm(torch.autograd.Function):
|
|||
r = torch.empty(n_rows, dtype = torch.float32, device = device)
|
||||
mu = torch.empty(n_rows, dtype = torch.float32, device = device)
|
||||
|
||||
layernorm_forward[(n_rows,)](
|
||||
Y, Y.stride(0),
|
||||
X, X.stride(0),
|
||||
W,
|
||||
b,
|
||||
r,
|
||||
mu,
|
||||
n_cols, eps,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
with torch_cuda_device(device):
|
||||
layernorm_forward[(n_rows,)](
|
||||
Y, Y.stride(0),
|
||||
X, X.stride(0),
|
||||
W,
|
||||
b,
|
||||
r,
|
||||
mu,
|
||||
n_cols, eps,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
ctx.eps = eps
|
||||
ctx.BLOCK_SIZE = BLOCK_SIZE
|
||||
ctx.num_warps = num_warps
|
||||
|
|
@ -137,17 +138,18 @@ class Fast_Layernorm(torch.autograd.Function):
|
|||
X, W, b, r, mu = ctx.saved_tensors
|
||||
n_rows, n_cols = dY.shape
|
||||
|
||||
layernorm_backward[(n_rows,)](
|
||||
dY, dY.stride(0),
|
||||
X, X .stride(0),
|
||||
W,
|
||||
b,
|
||||
r,
|
||||
mu,
|
||||
n_cols, ctx.eps,
|
||||
BLOCK_SIZE = ctx.BLOCK_SIZE,
|
||||
num_warps = ctx.num_warps,
|
||||
)
|
||||
with torch_cuda_device(dY.device):
|
||||
layernorm_backward[(n_rows,)](
|
||||
dY, dY.stride(0),
|
||||
X, X .stride(0),
|
||||
W,
|
||||
b,
|
||||
r,
|
||||
mu,
|
||||
n_cols, ctx.eps,
|
||||
BLOCK_SIZE = ctx.BLOCK_SIZE,
|
||||
num_warps = ctx.num_warps,
|
||||
)
|
||||
dX = dY.view(*shape)
|
||||
return dX, None, None, None, None
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -15,8 +15,7 @@
|
|||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
from .utils import calculate_settings
|
||||
|
||||
from .utils import calculate_settings, torch_cuda_device
|
||||
|
||||
@triton.jit
|
||||
def _rms_layernorm_forward(
|
||||
|
|
@ -154,15 +153,16 @@ class Fast_RMS_Layernorm(torch.autograd.Function):
|
|||
r = torch.empty(n_rows, dtype = torch.float32, device = device)
|
||||
|
||||
fx = _gemma_rms_layernorm_forward if gemma else _rms_layernorm_forward
|
||||
fx[(n_rows,)](
|
||||
Y, Y.stride(0),
|
||||
X, X.stride(0),
|
||||
W, W.stride(0),
|
||||
r, r.stride(0),
|
||||
n_cols, eps,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
with torch_cuda_device(device):
|
||||
fx[(n_rows,)](
|
||||
Y, Y.stride(0),
|
||||
X, X.stride(0),
|
||||
W, W.stride(0),
|
||||
r, r.stride(0),
|
||||
n_cols, eps,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
ctx.eps = eps
|
||||
ctx.BLOCK_SIZE = BLOCK_SIZE
|
||||
ctx.num_warps = num_warps
|
||||
|
|
@ -183,18 +183,19 @@ class Fast_RMS_Layernorm(torch.autograd.Function):
|
|||
# dW = X
|
||||
dX = torch.empty_like(dY) if ctx.GEMMA else dY
|
||||
|
||||
_rms_layernorm_backward[(n_rows,)](
|
||||
dY, dY.stride(0),
|
||||
dX, dX.stride(0),
|
||||
X, X .stride(0),
|
||||
W, W .stride(0),
|
||||
r, r .stride(0),
|
||||
# dW, dW.stride(0),
|
||||
n_cols, ctx.eps,
|
||||
GEMMA = ctx.GEMMA,
|
||||
BLOCK_SIZE = ctx.BLOCK_SIZE,
|
||||
num_warps = ctx.num_warps,
|
||||
)
|
||||
with torch_cuda_device(dY.device):
|
||||
_rms_layernorm_backward[(n_rows,)](
|
||||
dY, dY.stride(0),
|
||||
dX, dX.stride(0),
|
||||
X, X .stride(0),
|
||||
W, W .stride(0),
|
||||
r, r .stride(0),
|
||||
# dW, dW.stride(0),
|
||||
n_cols, ctx.eps,
|
||||
GEMMA = ctx.GEMMA,
|
||||
BLOCK_SIZE = ctx.BLOCK_SIZE,
|
||||
num_warps = ctx.num_warps,
|
||||
)
|
||||
dX = dX.view(*shape)
|
||||
return dX, None, None, None
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@
|
|||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
from .utils import calculate_settings
|
||||
from .utils import calculate_settings, torch_cuda_device
|
||||
ROPE_GROUP_SIZE : int = 4
|
||||
|
||||
def _rope_embedding(
|
||||
|
|
@ -100,16 +100,17 @@ class Fast_RoPE_Embedding(torch.autograd.Function):
|
|||
div, mod = divmod(n_heads, ROPE_GROUP_SIZE)
|
||||
n_groups : int = div + (mod != 0)
|
||||
|
||||
_rope_embedding[(n_rows, n_groups, )](
|
||||
Q, Q.stride(0),
|
||||
cos, cos.stride(0),
|
||||
sin, sin.stride(0),
|
||||
seq_len,
|
||||
head_dim, n_heads,
|
||||
BACKWARD_PASS = False,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
with torch_cuda_device(Q.device):
|
||||
_rope_embedding[(n_rows, n_groups, )](
|
||||
Q, Q.stride(0),
|
||||
cos, cos.stride(0),
|
||||
sin, sin.stride(0),
|
||||
seq_len,
|
||||
head_dim, n_heads,
|
||||
BACKWARD_PASS = False,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
ctx.BLOCK_SIZE = BLOCK_SIZE
|
||||
ctx.num_warps = num_warps
|
||||
ctx.n_groups = n_groups
|
||||
|
|
@ -134,15 +135,16 @@ class Fast_RoPE_Embedding(torch.autograd.Function):
|
|||
cos = ctx.cos
|
||||
sin = ctx.sin
|
||||
|
||||
_rope_embedding[(n_rows, ctx.n_groups, )](
|
||||
dY, dY .stride(0),
|
||||
cos, cos.stride(0),
|
||||
sin, sin.stride(0),
|
||||
seq_len, head_dim, n_heads,
|
||||
BACKWARD_PASS = True,
|
||||
BLOCK_SIZE = ctx.BLOCK_SIZE,
|
||||
num_warps = ctx.num_warps,
|
||||
)
|
||||
with torch_cuda_device(dY.device):
|
||||
_rope_embedding[(n_rows, ctx.n_groups, )](
|
||||
dY, dY .stride(0),
|
||||
cos, cos.stride(0),
|
||||
sin, sin.stride(0),
|
||||
seq_len, head_dim, n_heads,
|
||||
BACKWARD_PASS = True,
|
||||
BLOCK_SIZE = ctx.BLOCK_SIZE,
|
||||
num_warps = ctx.num_warps,
|
||||
)
|
||||
dY = dY.view(batch, seq_len, n_heads, head_dim)
|
||||
return dY, None, None,
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@
|
|||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
from .utils import calculate_settings
|
||||
from .utils import calculate_settings, torch_cuda_device
|
||||
|
||||
|
||||
@triton.jit
|
||||
|
|
@ -43,7 +43,8 @@ def swiglu_fg_kernel(e, g):
|
|||
n_elements = e.numel()
|
||||
h = torch.empty((batch, seq_len, hd), dtype = e.dtype, device = e.device)
|
||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||
_fg_kernel[grid](e, g, h, n_elements, BLOCK_SIZE = 1024,)
|
||||
with torch_cuda_device(e.device):
|
||||
_fg_kernel[grid](e, g, h, n_elements, BLOCK_SIZE = 1024,)
|
||||
return h
|
||||
pass
|
||||
|
||||
|
|
@ -94,6 +95,7 @@ def swiglu_DWf_DW_dfg_kernel(DW, e, g):
|
|||
batch_seq_len, hd = e.shape
|
||||
n_elements = e.numel()
|
||||
grid = lambda meta: (triton.cdiv(n_elements, meta['BLOCK_SIZE']),)
|
||||
_DWf_DW_dfg_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
||||
with torch_cuda_device(e.device):
|
||||
_DWf_DW_dfg_kernel[grid](DW, e, g, n_elements, BLOCK_SIZE = 1024,)
|
||||
return DW, e, g
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ import functools
|
|||
|
||||
# torch.cuda.amp.custom_fwd is deprecated >= 2.4
|
||||
import torch
|
||||
torch_Tensor = torch.Tensor
|
||||
from packaging.version import Version
|
||||
if Version(torch.__version__) < Version("2.4.0"):
|
||||
torch_amp_custom_fwd = torch.cuda.amp.custom_fwd
|
||||
|
|
@ -67,6 +68,18 @@ import ctypes
|
|||
HAS_CUDA_STREAM = Version(bnb.__version__) > Version("0.43.3")
|
||||
get_ptr = bnb.functional.get_ptr
|
||||
|
||||
if torch.cuda.device_count() > 1:
|
||||
torch_cuda_device = torch.cuda.device
|
||||
else:
|
||||
from contextlib import nullcontext
|
||||
def torch_cuda_device(device): return nullcontext()
|
||||
pass
|
||||
_cuda_getCurrentRawStream = torch._C._cuda_getCurrentRawStream
|
||||
c_void_p = ctypes.c_void_p
|
||||
def _get_tensor_stream(tensor: torch_Tensor) -> c_void_p:
|
||||
return c_void_p(_cuda_getCurrentRawStream(tensor.device.index))
|
||||
pass
|
||||
|
||||
# Get array of CUDA streams and other buffers
|
||||
global CUDA_STREAMS
|
||||
global WEIGHT_BUFFERS
|
||||
|
|
@ -92,27 +105,29 @@ cdequantize_blockwise_bf16_nf4 = bnb.functional.lib.cdequantize_blockwise_bf16_
|
|||
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemm_4bit_inference_naive_fp16
|
||||
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemm_4bit_inference_naive_bf16
|
||||
|
||||
|
||||
def QUANT_STATE(W):
|
||||
return getattr(W, "quant_state", None)
|
||||
pass
|
||||
|
||||
def QUANT_STATE(W): return getattr(W, "quant_state", None)
|
||||
|
||||
def get_lora_parameters(proj):
|
||||
# For DPO or disabled adapters
|
||||
base_layer = (proj.base_layer if hasattr(proj, "base_layer") else proj)
|
||||
base_layer = getattr(proj, "base_layer", proj) # (proj.base_layer if hasattr(proj, "base_layer") else proj)
|
||||
W = base_layer.weight
|
||||
|
||||
if not hasattr(proj, "disable_adapters") or proj.disable_adapters or proj.merged:
|
||||
return W, QUANT_STATE(W), None, None, None
|
||||
# if not hasattr(proj, "disable_adapters") or proj.disable_adapters or proj.merged:
|
||||
if getattr(proj, "disable_adapters", True) or proj.merged:
|
||||
return W, getattr(W, "quant_state", None), None, None, None
|
||||
pass
|
||||
|
||||
active_adapter = proj.active_adapters[0] if \
|
||||
hasattr(proj, "active_adapters") else proj.active_adapter
|
||||
A = proj.lora_A [active_adapter].weight
|
||||
B = proj.lora_B [active_adapter].weight
|
||||
s = proj.scaling[active_adapter]
|
||||
return W, QUANT_STATE(W), A, B, s
|
||||
adapter = getattr(proj, "active_adapters", None)
|
||||
if adapter is None: adapter = getattr(proj, "active_adapter", ("default"))
|
||||
adapter = adapter[0]
|
||||
|
||||
return (
|
||||
W,
|
||||
getattr(W, "quant_state", None),
|
||||
proj.lora_A [adapter].weight,
|
||||
proj.lora_B [adapter].weight,
|
||||
proj.scaling[adapter],
|
||||
)
|
||||
pass
|
||||
|
||||
|
||||
|
|
@ -120,19 +135,24 @@ def get_lora_parameters_bias(proj):
|
|||
# For DPO or disabled adapters
|
||||
base_layer = getattr(proj, "base_layer", proj) # (proj.base_layer if hasattr(proj, "base_layer") else proj)
|
||||
W = base_layer.weight
|
||||
bias = base_layer.bias
|
||||
|
||||
# if not hasattr(proj, "disable_adapters") or proj.disable_adapters or proj.merged:
|
||||
if getattr(proj, "disable_adapters", True) or proj.merged:
|
||||
return W, QUANT_STATE(W), None, None, None, bias
|
||||
return W, getattr(W, "quant_state", None), None, None, None, bias
|
||||
pass
|
||||
|
||||
active_adapter = proj.active_adapters[0] if \
|
||||
getattr(proj, "active_adapters", ) else proj.active_adapter
|
||||
A = proj.lora_A [active_adapter].weight
|
||||
B = proj.lora_B [active_adapter].weight
|
||||
s = proj.scaling[active_adapter]
|
||||
return W, QUANT_STATE(W), A, B, s, bias
|
||||
adapter = getattr(proj, "active_adapters", None)
|
||||
if adapter is None: adapter = getattr(proj, "active_adapter", ("default"))
|
||||
adapter = adapter[0]
|
||||
|
||||
return (
|
||||
W,
|
||||
getattr(W, "quant_state", None),
|
||||
proj.lora_A [adapter].weight,
|
||||
proj.lora_B [adapter].weight,
|
||||
proj.scaling[adapter],
|
||||
base_layer.bias,
|
||||
)
|
||||
pass
|
||||
|
||||
if HAS_CUDA_STREAM:
|
||||
|
|
@ -193,18 +213,19 @@ if HAS_CUDA_STREAM:
|
|||
|
||||
# NF4 dequantization of statistics
|
||||
ptr_out_absmax = get_ptr(out_absmax)
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), ptr_out_absmax,
|
||||
ctypes_c_int(blocksize2), ctypes_c_int(n_elements_absmax), CUDA_STREAM,
|
||||
)
|
||||
out_absmax += offset
|
||||
|
||||
# Dequantize W
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \
|
||||
cdequantize_blockwise_bf16_nf4
|
||||
fx(get_ptr(None), get_ptr(W), ptr_out_absmax, get_ptr(out),
|
||||
ctypes_c_int(blocksize), ctypes_c_int(out.numel()), CUDA_STREAM,)
|
||||
with torch_cuda_device(device):
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), ptr_out_absmax,
|
||||
ctypes_c_int(blocksize2), ctypes_c_int(n_elements_absmax), CUDA_STREAM
|
||||
)
|
||||
out_absmax += offset
|
||||
|
||||
# Dequantize W
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \
|
||||
cdequantize_blockwise_bf16_nf4
|
||||
fx(get_ptr(None), get_ptr(W), ptr_out_absmax, get_ptr(out),
|
||||
ctypes_c_int(blocksize), ctypes_c_int(out.numel()), CUDA_STREAM,)
|
||||
pass
|
||||
# Careful returning transposed data
|
||||
is_transposed = (True if W.shape[0] == 1 else False)
|
||||
return out.t() if is_transposed else out
|
||||
|
|
@ -316,19 +337,21 @@ if HAS_CUDA_STREAM:
|
|||
ldc = ctypes_c_int32(ldc)
|
||||
|
||||
df = torch.empty(absmax.shape, dtype = torch.float32, device = device)
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df),
|
||||
ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), CUDA_STREAM,
|
||||
)
|
||||
df += offset
|
||||
absmax = df
|
||||
with torch_cuda_device(device):
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df),
|
||||
ctypes_c_int(blocksize2), ctypes_c_int(df.numel()), CUDA_STREAM,
|
||||
)
|
||||
df += offset
|
||||
absmax = df
|
||||
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \
|
||||
cgemm_4bit_inference_naive_bf16
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \
|
||||
cgemm_4bit_inference_naive_bf16
|
||||
|
||||
blocksize = ctypes_c_int32(blocksize)
|
||||
fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out),
|
||||
lda, ldb, ldc, blocksize, CUDA_STREAM,)
|
||||
blocksize = ctypes_c_int32(blocksize)
|
||||
fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out),
|
||||
lda, ldb, ldc, blocksize, CUDA_STREAM,)
|
||||
pass
|
||||
|
||||
return out
|
||||
pass
|
||||
|
|
@ -458,7 +481,6 @@ def matmul_lora(X, W, W_quant, A, B, s, out = None):
|
|||
else:
|
||||
reshape = False
|
||||
pass
|
||||
|
||||
out = torch_matmul(X, W, out = out)
|
||||
if W_quant is not None: del W
|
||||
|
||||
|
|
|
|||
|
|
@ -19,5 +19,5 @@ from .llama import FastLlamaModel
|
|||
from .mistral import FastMistralModel
|
||||
from .qwen2 import FastQwen2Model
|
||||
from .dpo import PatchDPOTrainer, PatchKTOTrainer
|
||||
from ._utils import is_bfloat16_supported
|
||||
from ._utils import is_bfloat16_supported, __version__
|
||||
from .rl import PatchFastRL, vLLMSamplingParams
|
||||
|
|
|
|||
|
|
@ -755,7 +755,8 @@ def offload_to_disk(W, model, name, temporary_location : str = "_unsloth_tempora
|
|||
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)
|
||||
# We must use weights_only = False due to pickling
|
||||
offloaded_W = torch.load(filename, map_location = "cpu", mmap = True, weights_only = False)
|
||||
offloaded_W._offloaded_file_location = filename
|
||||
return offloaded_W
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ import math
|
|||
from functools import partial
|
||||
from typing import Optional, Tuple, List, Union
|
||||
from ._utils import *
|
||||
from ._utils import patch_unsloth_smart_gradient_checkpointing
|
||||
from ._utils import __version__
|
||||
from torch.nn.functional import scaled_dot_product_attention
|
||||
from transformers import __version__ as transformers_version
|
||||
|
|
@ -758,14 +759,9 @@ def LlamaModel_fast_forward(
|
|||
|
||||
# Check checkpointing method
|
||||
gradient_checkpointing = False
|
||||
offloaded_gradient_checkpointing = False
|
||||
|
||||
if (self.gradient_checkpointing and self.training and not use_cache):
|
||||
|
||||
gradient_checkpointing = True
|
||||
|
||||
if output_attentions is False and hasattr(self, "_offloaded_gradient_checkpointing"):
|
||||
offloaded_gradient_checkpointing = True
|
||||
pass
|
||||
|
||||
# Gemma2 has alternating SWA and global attn
|
||||
|
|
@ -850,27 +846,12 @@ def LlamaModel_fast_forward(
|
|||
mask = self. GA_mask if use_static_mask else dynamic_GA_mask
|
||||
pass
|
||||
|
||||
if offloaded_gradient_checkpointing:
|
||||
hidden_states = Unsloth_Offloaded_Gradient_Checkpointer.apply(
|
||||
decoder_layer,
|
||||
hidden_states,
|
||||
mask,
|
||||
attention_mask,
|
||||
position_ids,
|
||||
past_key_values,
|
||||
output_attentions,
|
||||
use_cache,
|
||||
None,
|
||||
position_embeddings,
|
||||
)[0]
|
||||
|
||||
elif gradient_checkpointing:
|
||||
if gradient_checkpointing:
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs, past_key_value, output_attentions, padding_mask = padding_mask, position_embeddings = position_embeddings)
|
||||
return custom_forward
|
||||
pass
|
||||
|
||||
layer_outputs = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(decoder_layer),
|
||||
hidden_states,
|
||||
|
|
@ -1703,10 +1684,10 @@ class FastLlamaModel:
|
|||
|
||||
statistics = \
|
||||
f"==((====))== Unsloth {__version__}: Fast {model_patcher.__name__[4:-5]} patching. Transformers: {transformers_version}.\n"\
|
||||
f" {chr(92)}{chr(92)} /| GPU: {gpu_stats.name}. Max memory: {max_memory} GB. Platform: {platform_system}.\n"\
|
||||
f" {chr(92)}{chr(92)} /| {gpu_stats.name}. Num GPUs = {torch.cuda.device_count()}. Max memory: {max_memory} GB. Platform: {platform_system}.\n"\
|
||||
f"O^O/ {chr(92)}_/ {chr(92)} Torch: {torch.__version__}. CUDA: {gpu_stats.major}.{gpu_stats.minor}. CUDA Toolkit: {torch.version.cuda}. Triton: {triton_version}\n"\
|
||||
f"{chr(92)} / Bfloat16 = {str(SUPPORTS_BFLOAT16).upper()}. FA [Xformers = {xformers_version}. FA2 = {HAS_FLASH_ATTENTION}]\n"\
|
||||
f' "-____-" Free Apache license: http://github.com/unslothai/unsloth'
|
||||
f' "-____-" Free license: http://github.com/unslothai/unsloth'
|
||||
print(statistics)
|
||||
|
||||
# Warn about fast transfers
|
||||
|
|
@ -1898,11 +1879,11 @@ class FastLlamaModel:
|
|||
# Cannot use \\ since it will cause a SyntaxWarning in Python 3.12
|
||||
# Instead use chr(92) == \\
|
||||
debug_info = """debug_info = \\
|
||||
f"==((====))== Unsloth - 2x faster free finetuning | Num GPUs = {args.world_size}\\n"\\
|
||||
f" {chr(92)}{chr(92)} /| Num examples = {num_examples:,} | Num Epochs = {num_train_epochs:,}\\n"\\
|
||||
f"O^O/ {chr(92)}_/ {chr(92)} Batch size per device = {self._train_batch_size:,} | Gradient Accumulation steps = {args.gradient_accumulation_steps}\\n"\\
|
||||
f"{chr(92)} / 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):,}'
|
||||
f"==((====))== Unsloth - 2x faster free finetuning | Num GPUs used = {len(set(p.device for p in model.parameters()))}\\n"\\
|
||||
f" {chr(92)}{chr(92)} /| Num examples = {num_examples:,} | Num Epochs = {num_train_epochs:,} | Total steps = {max_steps:,}\\n"\\
|
||||
f"O^O/ {chr(92)}_/ {chr(92)} Batch size per device = {self._train_batch_size:,} | Gradient accumulation steps = {args.gradient_accumulation_steps}\\n"\\
|
||||
f"{chr(92)} / Data Parallel GPUs = {args.world_size} | Total batch size ({self._train_batch_size} x {args.gradient_accumulation_steps} x {args.world_size}) = {total_train_batch_size:,}\\n"\\
|
||||
f' "-____-" Trainable parameters = {get_model_param_count(model, trainable_only=True):,}/{get_model_param_count(model):,} ({get_model_param_count(model, trainable_only=True)/get_model_param_count(model)*100:.2f}% trained)'
|
||||
logger.warning(debug_info)
|
||||
import subprocess, re, gc
|
||||
for _ in range(3):
|
||||
|
|
@ -1989,9 +1970,14 @@ class FastLlamaModel:
|
|||
internal_model = model
|
||||
while hasattr(internal_model, "model"):
|
||||
internal_model._saved_temp_tokenizer = tokenizer
|
||||
# Also set is_loaded_in_8bit to disable incorrect DDP
|
||||
internal_model.is_loaded_in_8bit = True
|
||||
|
||||
internal_model = internal_model.model
|
||||
pass
|
||||
internal_model._saved_temp_tokenizer = tokenizer
|
||||
# Also set is_loaded_in_8bit to disable incorrect DDP
|
||||
internal_model.is_loaded_in_8bit = True
|
||||
|
||||
# For transformers > 4.47.1, we need to add rotary_emb to all attention layers
|
||||
if IS_ATTENTION_REFACTOR or hasattr(model.model, "rotary_emb"):
|
||||
|
|
@ -2034,6 +2020,9 @@ class FastLlamaModel:
|
|||
):
|
||||
transformers_set_seed(random_state)
|
||||
|
||||
if use_gradient_checkpointing == "unsloth":
|
||||
patch_unsloth_smart_gradient_checkpointing(dtype = model.get_input_embeddings().weight.dtype)
|
||||
|
||||
if type(r) is not int:
|
||||
raise TypeError(f"Unsloth: Rank of {str(r)} must be an integer.")
|
||||
if r <= 0:
|
||||
|
|
@ -2398,11 +2387,15 @@ class FastLlamaModel:
|
|||
if hasattr(internal_model, "_saved_temp_tokenizer"):
|
||||
internal_model._saved_temp_tokenizer.padding_side = "right"
|
||||
pass
|
||||
# Also set is_loaded_in_8bit to disable incorrect DDP
|
||||
internal_model.is_loaded_in_8bit = True
|
||||
internal_model = internal_model.model
|
||||
pass
|
||||
if hasattr(internal_model, "_saved_temp_tokenizer"):
|
||||
internal_model._saved_temp_tokenizer.padding_side = "right"
|
||||
pass
|
||||
# Also set is_loaded_in_8bit to disable incorrect DDP
|
||||
internal_model.is_loaded_in_8bit = True
|
||||
|
||||
# Clear deleted GPU items
|
||||
for _ in range(3):
|
||||
|
|
|
|||
|
|
@ -59,7 +59,15 @@ if SUPPORTS_GEMMA2:
|
|||
from .gemma2 import FastGemma2Model
|
||||
pass
|
||||
import torch
|
||||
|
||||
from ._utils import (
|
||||
patch_compiling_bitsandbytes,
|
||||
patch_model_and_tokenizer,
|
||||
prepare_model_for_kbit_training,
|
||||
patch_unsloth_smart_gradient_checkpointing,
|
||||
patch_compiled_autograd,
|
||||
process_vision_info,
|
||||
unsloth_compile_transformers,
|
||||
)
|
||||
|
||||
class FastLanguageModel(FastLlamaModel):
|
||||
@staticmethod
|
||||
|
|
@ -87,6 +95,10 @@ class FastLanguageModel(FastLlamaModel):
|
|||
*args, **kwargs,
|
||||
):
|
||||
if token is None: token = get_token()
|
||||
assert (dtype is None or dtype == torch.float16 or dtype == torch.bfloat16)
|
||||
|
||||
if use_gradient_checkpointing == "unsloth":
|
||||
patch_unsloth_smart_gradient_checkpointing(dtype = dtype)
|
||||
|
||||
if fast_inference:
|
||||
if importlib.util.find_spec("vllm") is None:
|
||||
|
|
@ -367,15 +379,6 @@ class FastLanguageModel(FastLlamaModel):
|
|||
pass
|
||||
|
||||
|
||||
from ._utils import (
|
||||
patch_compiling_bitsandbytes,
|
||||
patch_model_and_tokenizer,
|
||||
prepare_model_for_kbit_training,
|
||||
patch_unsloth_smart_gradient_checkpointing,
|
||||
patch_compiled_autograd,
|
||||
process_vision_info,
|
||||
unsloth_compile_transformers,
|
||||
)
|
||||
from ..kernels import (
|
||||
patch_loss_functions,
|
||||
post_patch_loss_function,
|
||||
|
|
@ -404,6 +407,7 @@ class FastVisionModel(FastBaseVisionModel):
|
|||
*args, **kwargs,
|
||||
):
|
||||
if token is None: token = get_token()
|
||||
assert (dtype is None or dtype == torch.float16 or dtype == torch.bfloat16)
|
||||
|
||||
patch_compiled_autograd()
|
||||
patch_compiling_bitsandbytes()
|
||||
|
|
|
|||
|
|
@ -495,7 +495,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
RLTrainer_source,
|
||||
f"trl.trainer.{trainer_file}",
|
||||
imports,
|
||||
overwrite = True,
|
||||
overwrite = False,
|
||||
)
|
||||
|
||||
# Patch Trainer
|
||||
|
|
|
|||
108
unsloth/save.py
108
unsloth/save.py
|
|
@ -17,6 +17,8 @@ from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
|
|||
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
|
||||
from peft.tuners.lora import Linear as Peft_Linear
|
||||
from typing import Optional, Callable, Union, List
|
||||
import sys
|
||||
import requests
|
||||
import torch
|
||||
import os
|
||||
import shutil
|
||||
|
|
@ -1613,6 +1615,112 @@ def create_ollama_modelfile(tokenizer, gguf_location):
|
|||
return modelfile
|
||||
pass
|
||||
|
||||
def create_ollama_model(
|
||||
username: str,
|
||||
model_name: str,
|
||||
tag: str,
|
||||
modelfile_path: str
|
||||
):
|
||||
try:
|
||||
init_check = subprocess.run(
|
||||
['curl', 'http://localhost:11434'], capture_output=True, text=True, timeout=3
|
||||
)
|
||||
if init_check.returncode == 0:
|
||||
print(init_check.stdout.strip())
|
||||
else:
|
||||
print("Ollama Server is not Running")
|
||||
except subprocess.TimeoutExpired:
|
||||
return "Ollama Request Timeout"
|
||||
|
||||
process = subprocess.Popen(
|
||||
['ollama', 'create', f'{username}/{model_name}:{tag}', '-f', f'{modelfile_path}'],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
bufsize=1,
|
||||
universal_newlines=True
|
||||
)
|
||||
|
||||
for line in iter(process.stdout.readline, ''):
|
||||
print(line, end='')
|
||||
sys.stdout.flush()
|
||||
|
||||
return_code = process.wait()
|
||||
|
||||
if return_code != 0:
|
||||
print(f"\nMODEL CREATED FAILED WITH RETURN CODE {return_code}")
|
||||
else:
|
||||
print("\nMODEL CREATED SUCCESSFULLY")
|
||||
pass
|
||||
|
||||
|
||||
def push_to_ollama_hub(username: str, model_name: str, tag: str):
|
||||
try:
|
||||
init_check = subprocess.run(
|
||||
['curl', 'http://localhost:11434'], capture_output=True, text=True, timeout=3
|
||||
)
|
||||
if init_check.returncode == 0:
|
||||
print(init_check.stdout.strip())
|
||||
else:
|
||||
print("Ollama Server is not Running")
|
||||
except subprocess.TimeoutExpired:
|
||||
return "Ollama Request Timeout"
|
||||
|
||||
process = subprocess.Popen(
|
||||
['ollama', 'push', f'{username}/{model_name}:{tag}'],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
bufsize=1,
|
||||
universal_newlines=True
|
||||
)
|
||||
|
||||
for line in iter(process.stdout.readline, ''):
|
||||
print(line, end='')
|
||||
sys.stdout.flush()
|
||||
|
||||
return_code = process.wait()
|
||||
|
||||
if return_code != 0:
|
||||
print(f"\nMODEL PUBLISHED FAILED WITH RETURN CODE {return_code}")
|
||||
else:
|
||||
print("\nMODEL PUBLISHED SUCCESSFULLY")
|
||||
|
||||
|
||||
def push_to_ollama(
|
||||
tokenizer,
|
||||
gguf_location,
|
||||
username: str,
|
||||
model_name: str,
|
||||
tag: str
|
||||
):
|
||||
model_file = create_ollama_modelfile(
|
||||
tokenizer=tokenizer,
|
||||
gguf_location=gguf_location
|
||||
)
|
||||
|
||||
with open(f"Modelfile_{model_name}", "w") as f:
|
||||
f.write(model_file)
|
||||
f.close()
|
||||
|
||||
create_ollama_model(
|
||||
username=username,
|
||||
model_name=model_name,
|
||||
tag=tag,
|
||||
modelfile_path=f"Modelfile_{model_name}"
|
||||
)
|
||||
|
||||
push_to_ollama_hub(
|
||||
username=username,
|
||||
model_name=model_name,
|
||||
tag=tag
|
||||
)
|
||||
|
||||
print("Succesfully pushed to ollama")
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def unsloth_save_pretrained_gguf(
|
||||
self,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue