From e773a4ffe6b72c9ac9054d4735d47e4d69a8db8f Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 3 Nov 2024 18:35:43 -0800 Subject: [PATCH] Update cross_entropy_loss.py --- unsloth/kernels/cross_entropy_loss.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index a8af945221..04a2e1861c 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -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 ,