diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index 347f9eb9da..9c3c9442c9 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -32,16 +32,16 @@ from unsloth_zoo.loss_utils import ( @triton.jit def _cross_entropy_forward( logits_ptr , - logits_row_stride : tl.constexpr(tl.int64), + logits_row_stride , loss_ptr , logsumexp_ptr , labels_ptr , - VOCAB_SIZE : tl.constexpr(tl.int32), - BLOCK_SIZE : tl.constexpr(tl.int32), - DO_SOFTCAPPING : tl.constexpr(tl.int1), - SOFTCAP : tl.constexpr(tl.float32), - DO_LOGIT_SCALING : tl.constexpr(tl.int1), - LOGIT_SCALE : tl.constexpr(tl.float32), + VOCAB_SIZE , + BLOCK_SIZE : tl.constexpr, + DO_SOFTCAPPING , + SOFTCAP , + DO_LOGIT_SCALING , + LOGIT_SCALE , ): """ Cross Entropy Loss = 1/n sum [ -yi log(Pi) ] @@ -105,17 +105,17 @@ pass @triton.jit def _chunked_cross_entropy_forward( logits_ptr , - logits_row_stride : tl.constexpr(tl.int64), + logits_row_stride , loss_ptr , logsumexp_ptr , labels_ptr , - VOCAB_SIZE : tl.constexpr(tl.int32), - N_CHUNKS : tl.constexpr(tl.int32), - BLOCK_SIZE : tl.constexpr(tl.int32), - DO_SOFTCAPPING : tl.constexpr(tl.int1), - SOFTCAP : tl.constexpr(tl.float32), - DO_LOGIT_SCALING : tl.constexpr(tl.int1), - LOGIT_SCALE : tl.constexpr(tl.float32), + VOCAB_SIZE , + N_CHUNKS , + BLOCK_SIZE : tl.constexpr, + DO_SOFTCAPPING , + SOFTCAP , + DO_LOGIT_SCALING , + LOGIT_SCALE , ): """ 256K vocab divided in 4 chunks @@ -188,17 +188,17 @@ pass @triton.jit def _cross_entropy_backward( logits_ptr , - logits_row_stride : tl.constexpr(tl.int64), + logits_row_stride , dloss_ptr , dloss_row_stride , logsumexp_ptr , labels_ptr , - VOCAB_SIZE : tl.constexpr(tl.int32), - BLOCK_SIZE : tl.constexpr(tl.int32), - DO_SOFTCAPPING : tl.constexpr(tl.int1), - SOFTCAP : tl.constexpr(tl.float32), - DO_LOGIT_SCALING : tl.constexpr(tl.int1), - LOGIT_SCALE : tl.constexpr(tl.float32), + VOCAB_SIZE , + BLOCK_SIZE : tl.constexpr, + DO_SOFTCAPPING , + SOFTCAP , + DO_LOGIT_SCALING , + LOGIT_SCALE , ): """ CE_i = -y log(P) = y * (log[sum(exp(x))] - x)