Update cross_entropy_loss.py
This commit is contained in:
parent
7cf69dd9f2
commit
54ff6eb169
1 changed files with 0 additions and 17 deletions
|
|
@ -260,7 +260,6 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
pass
|
||||
|
||||
|
||||
# slow_cross_entropy_loss = torch.nn.functional.cross_entropy
|
||||
def fast_cross_entropy_loss(logits, labels):
|
||||
"""
|
||||
Arguments:
|
||||
|
|
@ -272,22 +271,6 @@ def fast_cross_entropy_loss(logits, labels):
|
|||
batch, seq_len, d = logits.shape
|
||||
assert(labels.shape == (batch, seq_len))
|
||||
|
||||
# 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),
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue