Fast CE Loss
This commit is contained in:
parent
25339b71f7
commit
603c71c7f0
4 changed files with 150 additions and 68 deletions
20
README.md
20
README.md
|
|
@ -113,8 +113,8 @@ pip install --upgrade --force-reinstall --no-cache-dir torch==2.1.0 triton \
|
|||
```bash
|
||||
pip install "unsloth[cu118] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118_ampere] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121_ampere] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118-ampere] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121-ampere] @ git+https://github.com/unslothai/unsloth.git"
|
||||
```
|
||||
3. For Pytorch 2.1.1: Use the `"ampere"` path for newer RTX 30xx GPUs or higher.
|
||||
```bash
|
||||
|
|
@ -122,10 +122,10 @@ pip install --upgrade --force-reinstall --no-cache-dir torch==2.1.1 triton \
|
|||
--index-url https://download.pytorch.org/whl/cu121
|
||||
```
|
||||
```bash
|
||||
pip install "unsloth[cu118_torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121_torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118_ampere_torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121_ampere_torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118-torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121-torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118-ampere-torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121-ampere-torch211] @ git+https://github.com/unslothai/unsloth.git"
|
||||
```
|
||||
4. For Pytorch 2.2.0: Use the `"ampere"` path for newer RTX 30xx GPUs or higher.
|
||||
```bash
|
||||
|
|
@ -133,10 +133,10 @@ pip install --upgrade --force-reinstall --no-cache-dir torch==2.2.0 triton \
|
|||
--index-url https://download.pytorch.org/whl/cu121
|
||||
```
|
||||
```bash
|
||||
pip install "unsloth[cu118_torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121_torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118_ampere_torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121_ampere_torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118-torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121-torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu118-ampere-torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
pip install "unsloth[cu121-ampere-torch220] @ git+https://github.com/unslothai/unsloth.git"
|
||||
```
|
||||
5. If you get errors, try the below first, then go back to step 1:
|
||||
```bash
|
||||
|
|
|
|||
|
|
@ -20,12 +20,14 @@ from transformers.models.llama.modeling_llama import logger
|
|||
|
||||
|
||||
@triton.jit
|
||||
def _cross_entropy_forward(logits_ptr, logits_row_stride,
|
||||
loss_ptr,
|
||||
lse_ptr,
|
||||
labels_ptr,
|
||||
n_cols,
|
||||
BLOCK_SIZE: tl.constexpr,):
|
||||
def _small_cross_entropy_forward(
|
||||
logits_ptr, logits_row_stride,
|
||||
loss_ptr,
|
||||
lse_ptr,
|
||||
labels_ptr,
|
||||
n_cols,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
Cross Entropy Loss = 1/n sum [ -yi log(Pi) ]
|
||||
Pi = exp(xi) / sum(exp(xi))
|
||||
|
|
@ -62,12 +64,64 @@ pass
|
|||
|
||||
|
||||
@triton.jit
|
||||
def _cross_entropy_backward(logits_ptr, logits_row_stride,
|
||||
dloss_ptr, dloss_row_stride,
|
||||
lse_ptr,
|
||||
labels_ptr,
|
||||
n_cols,
|
||||
BLOCK_SIZE: tl.constexpr,):
|
||||
def _large_cross_entropy_forward(
|
||||
logits_ptr, logits_row_stride,
|
||||
loss_ptr,
|
||||
lse_ptr,
|
||||
labels_ptr,
|
||||
n_rows,
|
||||
n_cols,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
Cross Entropy Loss = 1/n sum [ -yi log(Pi) ]
|
||||
Pi = exp(xi) / sum(exp(xi))
|
||||
CE_i = -y log(p) = -y log[ exp(x) / sum(exp(x)) ]
|
||||
= -y [ x - log[sum(exp(x))] ]
|
||||
= y * (log[sum(exp(x))] - x)
|
||||
If y == 0: CE_i = 0
|
||||
If y == 1: CE_i = logsumexp - x
|
||||
"""
|
||||
row_idx = tl.program_id(0)
|
||||
col_idx = tl.program_id(1)
|
||||
logits_ptr += row_idx * logits_row_stride
|
||||
loss_ptr += row_idx + col_idx*n_rows
|
||||
lse_ptr += row_idx + col_idx*n_rows
|
||||
labels_ptr += row_idx
|
||||
|
||||
col_offsets = col_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||
mask = col_offsets < n_cols
|
||||
|
||||
# Get labels and logits
|
||||
label_idx = tl.load(labels_ptr).to(tl.int64)
|
||||
logits = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")).to(tl.float32)
|
||||
max_logits = tl.max(logits, 0)
|
||||
# Maximum stops overflow
|
||||
lse = tl.log(tl.sum(tl.exp(logits - max_logits), 0)) + max_logits
|
||||
tl.store(lse_ptr, lse)
|
||||
|
||||
if (label_idx != -100) and \
|
||||
(label_idx >= (col_idx+0)*BLOCK_SIZE) and \
|
||||
(label_idx < min((col_idx+1)*BLOCK_SIZE, n_cols)):
|
||||
|
||||
loss = tl.load(logits_ptr + label_idx).to(tl.float32)
|
||||
lse = 0.0
|
||||
loss = lse - logits_label # We add the final logsumexp after a reduction
|
||||
else:
|
||||
loss = 0.0
|
||||
tl.store(loss_ptr, loss)
|
||||
pass
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _cross_entropy_backward(
|
||||
logits_ptr, logits_row_stride,
|
||||
dloss_ptr, dloss_row_stride,
|
||||
lse_ptr,
|
||||
labels_ptr,
|
||||
n_cols,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
CE_i = -y log(P) = y * (log[sum(exp(x))] - x)
|
||||
dC/dx = d/dx (y * log[sum(exp(x))] - x * y)
|
||||
|
|
@ -84,18 +138,18 @@ def _cross_entropy_backward(logits_ptr, logits_row_stride,
|
|||
If y == 1 and x != label: dC/dx = exp[x - logsumexp]
|
||||
"""
|
||||
row_idx = tl.program_id(0)
|
||||
logits_ptr += row_idx * logits_row_stride
|
||||
col_idx = tl.program_id(1)
|
||||
logits_ptr += row_idx * logits_row_stride.to(tl.int64)
|
||||
dloss_ptr += row_idx * dloss_row_stride
|
||||
col_offsets = tl.arange(0, BLOCK_SIZE)
|
||||
col_offsets = col_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||
mask = col_offsets < n_cols
|
||||
# TODO: Fixup int32 locations to int64
|
||||
label_idx = tl.load(labels_ptr + row_idx).to(tl.int32)
|
||||
|
||||
if label_idx != -100:
|
||||
dloss = tl.load(dloss_ptr)
|
||||
else:
|
||||
dloss = 0.0
|
||||
logits = tl.load(logits_ptr + col_offsets, mask = mask, other = 0).to(tl.float32)
|
||||
logits = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")).to(tl.float32)
|
||||
lse = tl.load(lse_ptr + row_idx)
|
||||
probs = tl.exp(logits - lse)
|
||||
|
||||
|
|
@ -104,26 +158,52 @@ def _cross_entropy_backward(logits_ptr, logits_row_stride,
|
|||
pass
|
||||
|
||||
|
||||
MAX_FUSED_SIZE = 65536 # 2**16
|
||||
|
||||
class Fast_CrossEntropyLoss(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, logits, labels):
|
||||
n_rows, n_cols = logits.shape
|
||||
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
||||
losses = torch.empty(n_rows, dtype = torch.float32, device = "cuda")
|
||||
logsumexp = torch.empty(n_rows, dtype = torch.float32, device = "cuda")
|
||||
|
||||
_cross_entropy_forward[(n_rows,)](
|
||||
logits, logits.stride(0),
|
||||
losses,
|
||||
logsumexp,
|
||||
labels,
|
||||
n_cols,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
div, mod = divmod(n_cols, MAX_FUSED_SIZE)
|
||||
n_splits = div + (mod != 0)
|
||||
|
||||
if n_splits == 1:
|
||||
# For small vocabs <= 65336 like Llama, Mistral
|
||||
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
||||
losses = torch.empty(n_rows, dtype = torch.float32, device = "cuda")
|
||||
logsumexp = torch.empty(n_rows, dtype = torch.float32, device = "cuda")
|
||||
|
||||
_small_cross_entropy_forward[(n_rows,)](
|
||||
logits, logits.stride(0),
|
||||
losses,
|
||||
logsumexp,
|
||||
labels,
|
||||
n_cols,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
else:
|
||||
# For small vocabs > 65336 like Gemma
|
||||
losses = torch.empty((n_rows, n_splits), dtype = torch.float32, device = "cuda")
|
||||
logsumexp = torch.empty((n_rows, n_splits), dtype = torch.float32, device = "cuda")
|
||||
|
||||
_large_cross_entropy_forward[(n_rows, n_splits,)](
|
||||
logits, logits.stride(0),
|
||||
losses,
|
||||
logsumexp,
|
||||
labels,
|
||||
n_rows,
|
||||
n_cols,
|
||||
BLOCK_SIZE = MAX_FUSED_SIZE,
|
||||
num_warps = 32,
|
||||
)
|
||||
logsumexp = torch.logsumexp(logsumexp, dim = 0) # Row sum
|
||||
losses = losses.sum(dim = 0) # Row sum
|
||||
losses += logsumexp # loss = lse - logits_label
|
||||
losses.masked_fill_(labels == -100, 0) # Padding tokens
|
||||
pass
|
||||
|
||||
ctx.BLOCK_SIZE = BLOCK_SIZE
|
||||
ctx.num_warps = num_warps
|
||||
ctx.save_for_backward(logits, logsumexp, labels)
|
||||
return losses
|
||||
pass
|
||||
|
|
@ -132,22 +212,23 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
def backward(ctx, dlosses):
|
||||
logits, logsumexp, labels = ctx.saved_tensors
|
||||
n_rows, n_cols = logits.shape
|
||||
grid = lambda meta: (n_rows, triton.cdiv(n_cols, meta["BLOCK_SIZE"]))
|
||||
|
||||
_cross_entropy_backward[(n_rows,)](
|
||||
_cross_entropy_backward[grid](
|
||||
logits, logits.stride(0),
|
||||
dlosses, dlosses.stride(0),
|
||||
logsumexp,
|
||||
labels,
|
||||
n_cols,
|
||||
BLOCK_SIZE = ctx.BLOCK_SIZE,
|
||||
num_warps = ctx.num_warps,
|
||||
BLOCK_SIZE = 4096,
|
||||
num_warps = 8,
|
||||
)
|
||||
return logits, None, None,
|
||||
pass
|
||||
pass
|
||||
|
||||
|
||||
slow_cross_entropy_loss = torch.nn.functional.cross_entropy
|
||||
# slow_cross_entropy_loss = torch.nn.functional.cross_entropy
|
||||
def fast_cross_entropy_loss(logits, labels):
|
||||
"""
|
||||
Arguments:
|
||||
|
|
@ -159,25 +240,26 @@ def fast_cross_entropy_loss(logits, labels):
|
|||
batch, seq_len, d = logits.shape
|
||||
assert(labels.shape == (batch, seq_len))
|
||||
|
||||
# We now support any vocab size due to Gemma!
|
||||
|
||||
# Prelim support Qwen, Deepseek other large vocab sizes > 2^16
|
||||
if d > MAX_FUSED_SIZE:
|
||||
logger.warning_once(
|
||||
f"Unsloth: Vocab size of {d} exceeds the max CUDA blocksize of {MAX_FUSED_SIZE}.\n"\
|
||||
"For now, Unsloth will use Pytorch's CrossEntropyLoss, which will entail a\n"\
|
||||
"25% increase in memory usage and be slower. Make an issue on \n"\
|
||||
"Unsloth's Github page if you want a faster and more memory efficient kernel!"
|
||||
)
|
||||
loss = slow_cross_entropy_loss(
|
||||
logits.float().view(batch*seq_len, d), # Must cast to float32 for numerical stability
|
||||
labels.view(-1),
|
||||
)
|
||||
return loss
|
||||
else:
|
||||
loss = Fast_CrossEntropyLoss.apply(
|
||||
logits.view(batch*seq_len, d),
|
||||
labels.view(-1),
|
||||
)
|
||||
n_items = torch.count_nonzero(labels != -100)
|
||||
return loss.sum() / n_items
|
||||
pass
|
||||
# if d > MAX_FUSED_SIZE:
|
||||
# logger.warning_once(
|
||||
# f"Unsloth: Vocab size of {d} exceeds the max CUDA blocksize of {MAX_FUSED_SIZE}.\n"\
|
||||
# "For now, Unsloth will use Pytorch's CrossEntropyLoss, which will entail a\n"\
|
||||
# "25% increase in memory usage and be slower. Make an issue on \n"\
|
||||
# "Unsloth's Github page if you want a faster and more memory efficient kernel!"
|
||||
# )
|
||||
# loss = slow_cross_entropy_loss(
|
||||
# logits.float().view(batch*seq_len, d), # Must cast to float32 for numerical stability
|
||||
# labels.view(-1),
|
||||
# )
|
||||
# return loss
|
||||
# else:
|
||||
loss = Fast_CrossEntropyLoss.apply(
|
||||
logits.view(batch*seq_len, d),
|
||||
labels.view(-1),
|
||||
)
|
||||
n_items = torch.count_nonzero(labels != -100)
|
||||
return loss.sum() / n_items
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -149,10 +149,10 @@ class LoRA_MLP(torch.autograd.Function):
|
|||
# dX += matmul_lora(de, gateW.t(), gateW_quant, gateB, gateA, gateS)
|
||||
upW = fast_dequantize(upW.t(), upW_quant)
|
||||
dX = torch.matmul(df, upW.t(), out = X)
|
||||
#del upW
|
||||
del upW
|
||||
dX += df @ upB.to(dtype).t() @ (upS * upA.to(dtype).t())
|
||||
|
||||
gateW = fast_dequantize(gateW.t(), gateW_quant, out = upW.t())
|
||||
gateW = fast_dequantize(gateW.t(), gateW_quant)
|
||||
dX += de @ gateW.t()
|
||||
del gateW
|
||||
dX += de @ gateB.to(dtype).t() @ (gateS * gateA.to(dtype).t())
|
||||
|
|
|
|||
|
|
@ -1399,7 +1399,7 @@ class FastLlamaModel:
|
|||
(down_proj.base_layer if hasattr(down_proj, "base_layer") else down_proj).bias is None:
|
||||
|
||||
# https://stackoverflow.com/questions/50599045/python-replacing-a-function-within-a-class-of-a-module
|
||||
# layer.mlp.forward = types.MethodType(apply_lora_mlp, layer.mlp)
|
||||
layer.mlp.forward = types.MethodType(apply_lora_mlp, layer.mlp)
|
||||
n_mlp += 1
|
||||
else:
|
||||
logger.warning_once(
|
||||
|
|
@ -1433,7 +1433,7 @@ class FastLlamaModel:
|
|||
if hasattr(o_proj, "lora_A") and \
|
||||
(o_proj.base_layer if hasattr(o_proj, "base_layer") else o_proj).bias is None:
|
||||
|
||||
# layer.self_attn.apply_o = apply_lora_o
|
||||
layer.self_attn.apply_o = apply_lora_o
|
||||
n_o += 1
|
||||
else:
|
||||
logger.warning_once(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue