From ba19344fb9400a7f0c0d2887977a74f81d3dead8 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sat, 24 Feb 2024 18:03:45 +1100 Subject: [PATCH] Update cross_entropy_loss.py --- unsloth/kernels/cross_entropy_loss.py | 38 +++++++++++++-------------- 1 file changed, 19 insertions(+), 19 deletions(-) diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index 07af99e87f..2fd95ecfd4 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -246,23 +246,23 @@ def fast_cross_entropy_loss(logits, labels): # 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 + 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