diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index 9c3c9442c9..35fbbed730 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -37,7 +37,7 @@ def _cross_entropy_forward( logsumexp_ptr , labels_ptr , VOCAB_SIZE , - BLOCK_SIZE : tl.constexpr, + BLOCK_SIZE : tl.constexpr , DO_SOFTCAPPING , SOFTCAP , DO_LOGIT_SCALING , @@ -69,7 +69,7 @@ def _cross_entropy_forward( logsumexp_ptr += row_idx labels_ptr += row_idx - col_offsets = tl.arange(0, BLOCK_SIZE) + col_offsets = tl.arange(0, BLOCK_SIZE : tl.constexpr) mask = col_offsets < VOCAB_SIZE label_idx = tl.load(labels_ptr).to(tl.int32) @@ -111,7 +111,7 @@ def _chunked_cross_entropy_forward( labels_ptr , VOCAB_SIZE , N_CHUNKS , - BLOCK_SIZE : tl.constexpr, + BLOCK_SIZE : tl.constexpr , DO_SOFTCAPPING , SOFTCAP , DO_LOGIT_SCALING , @@ -148,7 +148,7 @@ def _chunked_cross_entropy_forward( logsumexp_ptr += row_idx * N_CHUNKS + chunk_idx labels_ptr += row_idx - col_offsets = chunk_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + col_offsets = chunk_idx*BLOCK_SIZE : tl.constexpr + tl.arange(0, BLOCK_SIZE : tl.constexpr) mask = col_offsets < VOCAB_SIZE label_idx = tl.load(labels_ptr).to(tl.int32) @@ -194,7 +194,7 @@ def _cross_entropy_backward( logsumexp_ptr , labels_ptr , VOCAB_SIZE , - BLOCK_SIZE : tl.constexpr, + BLOCK_SIZE : tl.constexpr , DO_SOFTCAPPING , SOFTCAP , DO_LOGIT_SCALING , @@ -220,7 +220,7 @@ def _cross_entropy_backward( logits_ptr += row_idx * logits_row_stride dloss_ptr += row_idx * dloss_row_stride - col_offsets = block_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + col_offsets = block_idx*BLOCK_SIZE : tl.constexpr + tl.arange(0, BLOCK_SIZE : tl.constexpr) mask = col_offsets < VOCAB_SIZE label_idx = tl.load(labels_ptr + row_idx).to(tl.int32) @@ -279,9 +279,12 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): n_chunks = div + (mod != 0) losses = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0") + DO_SOFTCAPPING : bool = bool(logit_softcapping != 0) + DO_LOGIT_SCALING : bool = bool(logit_scaling != 0) + if n_chunks == 1: # For small vocabs <= 65336 like Llama, Mistral - BLOCK_SIZE, num_warps = calculate_settings(vocab_size) + BLOCK_SIZE : tl.constexpr, num_warps = calculate_settings(vocab_size) logsumexp = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0") _cross_entropy_forward[(n_rows,)]( @@ -290,10 +293,10 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): logsumexp, labels, VOCAB_SIZE = vocab_size, - BLOCK_SIZE = BLOCK_SIZE, - DO_SOFTCAPPING = bool(logit_softcapping != 0), + BLOCK_SIZE : tl.constexpr = BLOCK_SIZE : tl.constexpr, + DO_SOFTCAPPING = DO_SOFTCAPPING, SOFTCAP = logit_softcapping, - DO_LOGIT_SCALING = bool(logit_scaling != 0), + DO_LOGIT_SCALING = DO_LOGIT_SCALING, LOGIT_SCALE = logit_scaling, num_warps = num_warps, ) @@ -308,10 +311,10 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): labels, VOCAB_SIZE = vocab_size, N_CHUNKS = n_chunks, - BLOCK_SIZE = MAX_FUSED_SIZE, - DO_SOFTCAPPING = bool(logit_softcapping != 0), + BLOCK_SIZE : tl.constexpr = MAX_FUSED_SIZE, + DO_SOFTCAPPING = DO_SOFTCAPPING, SOFTCAP = logit_softcapping, - DO_LOGIT_SCALING = bool(logit_scaling != 0), + DO_LOGIT_SCALING = DO_LOGIT_SCALING, LOGIT_SCALE = logit_scaling, num_warps = 32, ) @@ -323,9 +326,9 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): pass ctx.save_for_backward(logits, logsumexp, labels) - ctx.DO_SOFTCAPPING = bool(logit_softcapping != 0) + ctx.DO_SOFTCAPPING = DO_SOFTCAPPING ctx.logit_softcapping = logit_softcapping - ctx.DO_LOGIT_SCALING = bool(logit_scaling != 0) + ctx.DO_LOGIT_SCALING = DO_LOGIT_SCALING ctx.logit_scaling = logit_scaling return losses pass @@ -335,8 +338,8 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): logits, logsumexp, labels = ctx.saved_tensors n_rows, vocab_size = logits.shape - BLOCK_SIZE = 4096 - div, mod = divmod(vocab_size, BLOCK_SIZE) + BLOCK_SIZE : tl.constexpr = 4096 + div, mod = divmod(vocab_size, BLOCK_SIZE : tl.constexpr) n_blocks = div + (mod != 0) _cross_entropy_backward[(n_rows, n_blocks,)]( @@ -345,10 +348,10 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): logsumexp, labels, VOCAB_SIZE = vocab_size, - BLOCK_SIZE = BLOCK_SIZE, - DO_SOFTCAPPING = bool(ctx.DO_SOFTCAPPING), + BLOCK_SIZE : tl.constexpr = BLOCK_SIZE : tl.constexpr, + DO_SOFTCAPPING = ctx.DO_SOFTCAPPING, SOFTCAP = ctx.logit_softcapping, - DO_LOGIT_SCALING = bool(ctx.DO_LOGIT_SCALING), + DO_LOGIT_SCALING = ctx.DO_LOGIT_SCALING, LOGIT_SCALE = ctx.logit_scaling, num_warps = 8, )