diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index 009e351daa..f17a560200 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