Update cross_entropy_loss.py

This commit is contained in:
Daniel Han-Chen 2024-02-26 03:53:49 +11:00
commit a5abe39ded

View file

@ -73,55 +73,65 @@ pass
@triton.jit
def _large_cross_entropy_forward(
def _chunked_cross_entropy_forward(
logits_ptr, logits_row_stride,
loss_ptr,
lse_ptr,
logsumexp_ptr,
labels_ptr,
n_rows,
n_cols,
BLOCK_SIZE: tl.constexpr,
VOCAB_SIZE : tl.constexpr,
N_CHUNKS : tl.constexpr,
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)
256K vocab divided in 4 chunks
|-65536-| |-65536-| |-65536-| |-65536-|
|-------| |-------| |-------| |-------|
|-------| |-------| |-------| |-------|
If y == 0: CE_i = 0
If y == 1: CE_i = logsumexp - x
Notice we can do logsumexp for each chunk and then
logsumexp[chunk_sum(logsumexp)] == logsumexp
chunk_sum = log[chunk_sum(logsumexp)]
= log[exp(logsumexp(a)) + ... + exp(logsumexp(z))]
= log[exp(log[sum(exp(a))]) + ... + exp(log[sum(exp(z))])]
= log[sum(exp(a)) + ... + sum(exp(z))]
= logsumexp(x)
This means we can perform a logsumexp for each chunk, then do a
final logsumexp reduction!
Ie do: logsumexp(chunked_logsumexp) - x
"""
row_idx = tl.program_id(0)
col_idx = tl.program_id(1)
logits_ptr += row_idx * logits_row_stride.to(tl.int64)
loss_ptr += row_idx + col_idx*n_rows
lse_ptr += row_idx + col_idx*n_rows
labels_ptr += row_idx
row_idx = tl.program_id(0)
chunk_idx = tl.program_id(1)
logits_ptr += row_idx * logits_row_stride.to(tl.int64)
loss_ptr += row_idx
logsumexp_ptr += row_idx * N_CHUNKS + chunk_idx
labels_ptr += row_idx
col_offsets = col_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = col_offsets < n_cols
col_offsets = chunk_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = col_offsets < VOCAB_SIZE
# Get labels and logits
label_idx = tl.load(labels_ptr).to(tl.int64)
label_idx = tl.load(labels_ptr).to(tl.int32)
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)
c = tl.max(logits, 0)
logsumexp = c + tl.log(tl.sum(tl.exp(logits - c), 0))
loss = 0.0
# chained boolean operators (A or B or C) are not supported; use parentheses to split the chain.
if (label_idx != -100):
if (label_idx >= (col_idx+0)*BLOCK_SIZE) and \
(label_idx < min((col_idx+1)*BLOCK_SIZE, n_cols)):
logits_label = tl.load(logits_ptr + label_idx).to(tl.float32)
lse = 0.0
loss = lse - logits_label # We add the final logsumexp after a reduction
pass
if chunk_idx == 0:
# logsumexp(chunked_logsumexp) - x
# Do the -x separately
if label_idx != -100:
x = tl.load(logits_ptr + label_idx).to(tl.float32)
loss = -1.0 * x
else:
loss = 0.0
tl.store(loss_ptr, loss)
pass
tl.store(loss_ptr, loss)
tl.store(logsumexp_ptr, logsumexp)
pass
@ -186,11 +196,11 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
div, mod = divmod(vocab_size, MAX_FUSED_SIZE)
n_chunks = div + (mod != 0)
losses = torch.empty(n_rows, dtype = torch.float32, device = "cuda")
if n_chunks == 1:
# For small vocabs <= 65336 like Llama, Mistral
BLOCK_SIZE, num_warps = calculate_settings(vocab_size)
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,)](
@ -204,23 +214,23 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
)
else:
# For large vocabs > 65336 like Gemma 256K
losses = torch.empty((n_chunks, n_rows), dtype = torch.float32, device = "cuda")
logsumexp = torch.empty((n_chunks, n_rows), dtype = torch.float32, device = "cuda")
logsumexp = torch.empty((n_rows, n_chunks,), dtype = torch.float32, device = "cuda")
_large_cross_entropy_forward[(n_rows, n_chunks,)](
_chunked_cross_entropy_forward[(n_rows, n_chunks,)](
logits, logits.stride(0),
losses,
logsumexp,
labels,
n_rows,
vocab_size,
VOCAB_SIZE = vocab_size,
N_CHUNKS = n_chunks,
BLOCK_SIZE = MAX_FUSED_SIZE,
num_warps = 32,
)
logsumexp = torch.logsumexp(logsumexp, dim = 0) # Column sum
losses = losses.sum(dim = 0) # Column sum
losses += logsumexp # loss = lse - logits_label
losses.masked_fill_(labels == -100, 0) # Padding tokens
# logsumexp(chunked_logsumexp) - x
# Do the -x separately
logsumexp = torch.logsumexp(logsumexp, dim = 1) # Row sum
losses += logsumexp
losses.masked_fill_(labels == -100, 0) # Don't forget to mask padding out!
pass
ctx.save_for_backward(logits, logsumexp, labels)