Fix Triton heuristics
https://github.com/triton-lang/triton/issues/5224
This commit is contained in:
parent
7d7a1b0ef4
commit
bfce3d402c
3 changed files with 33 additions and 20 deletions
|
|
@ -25,11 +25,6 @@ from unsloth_zoo.loss_utils import (
|
|||
)
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
"DO_SOFTCAPPING": lambda args: bool(args["DO_SOFTCAPPING" ]),
|
||||
"DO_LOGIT_SCALING": lambda args: bool(args["DO_LOGIT_SCALING"]),
|
||||
})
|
||||
@triton.jit
|
||||
def _cross_entropy_forward(
|
||||
logits_ptr ,
|
||||
logits_row_stride ,
|
||||
|
|
@ -95,13 +90,15 @@ def _cross_entropy_forward(
|
|||
tl.store(logsumexp_ptr, logsumexp)
|
||||
tl.store(loss_ptr, loss)
|
||||
pass
|
||||
_cross_entropy_forward = triton.jit(_cross_entropy_forward)
|
||||
_cross_entropy_forward = triton.heuristics(
|
||||
{
|
||||
"DO_SOFTCAPPING": lambda args: bool(args["DO_SOFTCAPPING" ]),
|
||||
"DO_LOGIT_SCALING": lambda args: bool(args["DO_LOGIT_SCALING"]),
|
||||
}
|
||||
)(_cross_entropy_forward)
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
"DO_SOFTCAPPING": lambda args: bool(args["DO_SOFTCAPPING" ]),
|
||||
"DO_LOGIT_SCALING": lambda args: bool(args["DO_LOGIT_SCALING"]),
|
||||
})
|
||||
@triton.jit
|
||||
def _chunked_cross_entropy_forward(
|
||||
logits_ptr ,
|
||||
logits_row_stride ,
|
||||
|
|
@ -177,13 +174,15 @@ def _chunked_cross_entropy_forward(
|
|||
pass
|
||||
tl.store(logsumexp_ptr, logsumexp)
|
||||
pass
|
||||
_chunked_cross_entropy_forward = triton.jit(_chunked_cross_entropy_forward)
|
||||
_chunked_cross_entropy_forward = triton.heuristics(
|
||||
{
|
||||
"DO_SOFTCAPPING": lambda args: bool(args["DO_SOFTCAPPING" ]),
|
||||
"DO_LOGIT_SCALING": lambda args: bool(args["DO_LOGIT_SCALING"]),
|
||||
}
|
||||
)(_chunked_cross_entropy_forward)
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
"DO_SOFTCAPPING": lambda args: bool(args["DO_SOFTCAPPING" ]),
|
||||
"DO_LOGIT_SCALING": lambda args: bool(args["DO_LOGIT_SCALING"]),
|
||||
})
|
||||
@triton.jit
|
||||
def _cross_entropy_backward(
|
||||
logits_ptr ,
|
||||
logits_row_stride ,
|
||||
|
|
@ -264,10 +263,16 @@ def _cross_entropy_backward(
|
|||
# If y == 0: dC/dx = 0 ==> we already masked it to be = 0, so dloss = 0.
|
||||
tl.store(logits_ptr + col_offsets, dloss * y, mask = mask)
|
||||
pass
|
||||
_cross_entropy_backward = triton.jit(_cross_entropy_backward)
|
||||
_cross_entropy_backward = triton.heuristics(
|
||||
{
|
||||
"DO_SOFTCAPPING": lambda args: bool(args["DO_SOFTCAPPING" ]),
|
||||
"DO_LOGIT_SCALING": lambda args: bool(args["DO_LOGIT_SCALING"]),
|
||||
}
|
||||
)(_cross_entropy_backward)
|
||||
|
||||
|
||||
MAX_FUSED_SIZE = 65536 # 2**16
|
||||
|
||||
class Fast_CrossEntropyLoss(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, logits, labels, logit_softcapping : float = 0, logit_scaling : float = 0):
|
||||
|
|
|
|||
|
|
@ -53,8 +53,6 @@ def _rms_layernorm_forward(
|
|||
pass
|
||||
|
||||
|
||||
@triton.heuristics({"GEMMA": lambda args: bool(args["GEMMA"]),})
|
||||
@triton.jit
|
||||
def _rms_layernorm_backward(
|
||||
dY, dY_row_stride,
|
||||
dX, dX_row_stride,
|
||||
|
|
@ -97,6 +95,12 @@ def _rms_layernorm_backward(
|
|||
output = inv_var/n_cols * (n_cols*dY_W - normed*rowsum_dY_normed)
|
||||
tl.store(dX + col_offsets, output, mask = mask)
|
||||
pass
|
||||
_rms_layernorm_backward = triton.jit(_rms_layernorm_backward)
|
||||
_rms_layernorm_backward = triton.heuristics(
|
||||
{
|
||||
"GEMMA": lambda args: bool(args["GEMMA"]),
|
||||
}
|
||||
)(_rms_layernorm_backward)
|
||||
|
||||
|
||||
@triton.jit
|
||||
|
|
|
|||
|
|
@ -18,8 +18,6 @@ import torch
|
|||
from .utils import calculate_settings
|
||||
ROPE_GROUP_SIZE : int = 4
|
||||
|
||||
@triton.heuristics({"BACKWARD_PASS": lambda args: bool(args["BACKWARD_PASS"]),})
|
||||
@triton.jit
|
||||
def _rope_embedding(
|
||||
Q, Q_row_stride,
|
||||
cos, cos_row_stride,
|
||||
|
|
@ -69,6 +67,12 @@ def _rope_embedding(
|
|||
tl.store(Q + offs_q2, Q2*cos1 + Q1*sin1, mask = mask)
|
||||
pass
|
||||
pass
|
||||
_rope_embedding = triton.jit(_rope_embedding)
|
||||
_rope_embedding = triton.heuristics(
|
||||
{
|
||||
"BACKWARD_PASS": lambda args: bool(args["BACKWARD_PASS"]),
|
||||
}
|
||||
)(_rope_embedding)
|
||||
|
||||
|
||||
class Fast_RoPE_Embedding(torch.autograd.Function):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue