Update cross_entropy_loss.py
This commit is contained in:
parent
e773a4ffe6
commit
2cff29f70c
1 changed files with 9 additions and 6 deletions
|
|
@ -279,6 +279,9 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
n_chunks = div + (mod != 0)
|
||||
losses = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0")
|
||||
|
||||
DO_SOFTCAPPING : bool = bool(logit_softcapping != 0)
|
||||
DO_LOGIT_SCALING : bool = bool(logit_scaling != 0)
|
||||
|
||||
if n_chunks == 1:
|
||||
# For small vocabs <= 65336 like Llama, Mistral
|
||||
BLOCK_SIZE, num_warps = calculate_settings(vocab_size)
|
||||
|
|
@ -291,9 +294,9 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
labels,
|
||||
VOCAB_SIZE = vocab_size,
|
||||
BLOCK_SIZE = BLOCK_SIZE,
|
||||
DO_SOFTCAPPING = logit_softcapping != 0,
|
||||
DO_SOFTCAPPING = DO_SOFTCAPPING,
|
||||
SOFTCAP = logit_softcapping,
|
||||
DO_LOGIT_SCALING = logit_scaling != 0,
|
||||
DO_LOGIT_SCALING = DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = logit_scaling,
|
||||
num_warps = num_warps,
|
||||
)
|
||||
|
|
@ -309,9 +312,9 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
VOCAB_SIZE = vocab_size,
|
||||
N_CHUNKS = n_chunks,
|
||||
BLOCK_SIZE = MAX_FUSED_SIZE,
|
||||
DO_SOFTCAPPING = logit_softcapping != 0,
|
||||
DO_SOFTCAPPING = DO_SOFTCAPPING,
|
||||
SOFTCAP = logit_softcapping,
|
||||
DO_LOGIT_SCALING = logit_scaling != 0,
|
||||
DO_LOGIT_SCALING = DO_LOGIT_SCALING,
|
||||
LOGIT_SCALE = logit_scaling,
|
||||
num_warps = 32,
|
||||
)
|
||||
|
|
@ -323,9 +326,9 @@ class Fast_CrossEntropyLoss(torch.autograd.Function):
|
|||
pass
|
||||
|
||||
ctx.save_for_backward(logits, logsumexp, labels)
|
||||
ctx.DO_SOFTCAPPING = logit_softcapping != 0
|
||||
ctx.DO_SOFTCAPPING = DO_SOFTCAPPING
|
||||
ctx.logit_softcapping = logit_softcapping
|
||||
ctx.DO_LOGIT_SCALING = logit_scaling != 0
|
||||
ctx.DO_LOGIT_SCALING = DO_LOGIT_SCALING
|
||||
ctx.logit_scaling = logit_scaling
|
||||
return losses
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue