Update cross_entropy_loss.py
This commit is contained in:
parent
f3a93319b0
commit
e773a4ffe6
1 changed files with 12 additions and 12 deletions
|
|
@ -25,10 +25,10 @@ from unsloth_zoo.loss_utils import (
|
|||
)
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
"DO_SOFTCAPPING": lambda args: args["DO_SOFTCAPPING" ],
|
||||
"DO_LOGIT_SCALING": lambda args: args["DO_LOGIT_SCALING"],
|
||||
})
|
||||
# @triton.heuristics({
|
||||
# "DO_SOFTCAPPING": lambda args: args["DO_SOFTCAPPING" ],
|
||||
# "DO_LOGIT_SCALING": lambda args: args["DO_LOGIT_SCALING"],
|
||||
# })
|
||||
@triton.jit
|
||||
def _cross_entropy_forward(
|
||||
logits_ptr ,
|
||||
|
|
@ -98,10 +98,10 @@ def _cross_entropy_forward(
|
|||
pass
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
"DO_SOFTCAPPING": lambda args: args["DO_SOFTCAPPING" ],
|
||||
"DO_LOGIT_SCALING": lambda args: args["DO_LOGIT_SCALING"],
|
||||
})
|
||||
# @triton.heuristics({
|
||||
# "DO_SOFTCAPPING": lambda args: args["DO_SOFTCAPPING" ],
|
||||
# "DO_LOGIT_SCALING": lambda args: args["DO_LOGIT_SCALING"],
|
||||
# })
|
||||
@triton.jit
|
||||
def _chunked_cross_entropy_forward(
|
||||
logits_ptr ,
|
||||
|
|
@ -181,10 +181,10 @@ def _chunked_cross_entropy_forward(
|
|||
pass
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
"DO_SOFTCAPPING": lambda args: args["DO_SOFTCAPPING" ],
|
||||
"DO_LOGIT_SCALING": lambda args: args["DO_LOGIT_SCALING"],
|
||||
})
|
||||
# @triton.heuristics({
|
||||
# "DO_SOFTCAPPING": lambda args: args["DO_SOFTCAPPING" ],
|
||||
# "DO_LOGIT_SCALING": lambda args: args["DO_LOGIT_SCALING"],
|
||||
# })
|
||||
@triton.jit
|
||||
def _cross_entropy_backward(
|
||||
logits_ptr ,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue