From dc552adcf11d1f3735cb5b13f5846ebcf4da3bcb Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 3 Nov 2024 21:03:24 -0800 Subject: [PATCH] Update cross_entropy_loss.py --- unsloth/kernels/cross_entropy_loss.py | 32 +++++++++++++-------------- 1 file changed, 16 insertions(+), 16 deletions(-) diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index 61a015d9ba..efe18f5c14 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 , + logits_row_stride : tl.constexpr(tl.int64), loss_ptr , logsumexp_ptr , labels_ptr , - VOCAB_SIZE , - BLOCK_SIZE : tl.constexpr, + VOCAB_SIZE : tl.constexpr(tl.int32), + BLOCK_SIZE : tl.constexpr(tl.int32), DO_SOFTCAPPING , - SOFTCAP , + SOFTCAP : tl.constexpr(tl.float32), DO_LOGIT_SCALING , - LOGIT_SCALE , + LOGIT_SCALE : tl.constexpr(tl.float32), ): """ 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 , + logits_row_stride : tl.constexpr(tl.int64), loss_ptr , logsumexp_ptr , labels_ptr , - VOCAB_SIZE , - N_CHUNKS , - BLOCK_SIZE : tl.constexpr, + VOCAB_SIZE : tl.constexpr(tl.int32), + N_CHUNKS : tl.constexpr(tl.int32), + BLOCK_SIZE : tl.constexpr(tl.int32), DO_SOFTCAPPING , - SOFTCAP , + SOFTCAP : tl.constexpr(tl.float32), DO_LOGIT_SCALING , - LOGIT_SCALE , + LOGIT_SCALE : tl.constexpr(tl.float32), ): """ 256K vocab divided in 4 chunks @@ -188,17 +188,17 @@ pass @triton.jit def _cross_entropy_backward( logits_ptr , - logits_row_stride , + logits_row_stride : tl.constexpr(tl.int64), dloss_ptr , dloss_row_stride , logsumexp_ptr , labels_ptr , - VOCAB_SIZE , - BLOCK_SIZE : tl.constexpr, + VOCAB_SIZE : tl.constexpr(tl.int32), + BLOCK_SIZE : tl.constexpr(tl.int32), DO_SOFTCAPPING , - SOFTCAP , + SOFTCAP : tl.constexpr(tl.float32), DO_LOGIT_SCALING , - LOGIT_SCALE , + LOGIT_SCALE : tl.constexpr(tl.float32), ): """ CE_i = -y log(P) = y * (log[sum(exp(x))] - x)