From d4557513032b582d9a4e4d9fc3efd9484f9705e8 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 31 Oct 2024 01:56:03 -0700 Subject: [PATCH] Update cross_entropy_loss.py --- unsloth/kernels/cross_entropy_loss.py | 64 +++++++++++++-------------- 1 file changed, 32 insertions(+), 32 deletions(-) diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index a13475e227..f256918746 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -19,23 +19,23 @@ from .utils import calculate_settings, MAX_FUSED_SIZE, triton_tanh from transformers.models.llama.modeling_llama import logger from packaging.version import Version -# @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 , - logits_row_stride , + logits_row_stride : tl.constexpr(tl.int64), loss_ptr , logsumexp_ptr , labels_ptr , - VOCAB_SIZE , + VOCAB_SIZE : tl.constexpr, BLOCK_SIZE : tl.constexpr, - DO_SOFTCAPPING , - SOFTCAP , - DO_LOGIT_SCALING , - LOGIT_SCALE , + DO_SOFTCAPPING : tl.constexpr(tl.int1), + SOFTCAP : tl.constexpr(tl.float32), + DO_LOGIT_SCALING : tl.constexpr(tl.int1), + LOGIT_SCALE : tl.constexpr(tl.float32), ): """ Cross Entropy Loss = 1/n sum [ -yi log(Pi) ] @@ -67,14 +67,14 @@ def _cross_entropy_forward( mask = col_offsets < VOCAB_SIZE label_idx = tl.load(labels_ptr).to(tl.int32) - logits = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")).to(tl.float32) + logits = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")) # Go logit scaling for Cohere: t * x if DO_LOGIT_SCALING: logits = LOGIT_SCALE * logits # Do logit softcapping for Gemma 2: t * tanh(1/t * x) - if DO_SOFTCAPPING: logits = SOFTCAP * triton_tanh(logits / SOFTCAP) + if DO_SOFTCAPPING: logits = SOFTCAP * triton_tanh(logits.to(tl.float32) / SOFTCAP).to(logits.dtype) - # logits = logits + logits = logits.to(tl.float32) c = tl.max(logits, 0) logsumexp = c + tl.log(tl.sum(tl.exp(logits - c), 0)) @@ -99,7 +99,7 @@ pass @triton.jit def _chunked_cross_entropy_forward( logits_ptr , - logits_row_stride , + logits_row_stride : tl.constexpr(tl.int64), loss_ptr , logsumexp_ptr , labels_ptr , @@ -107,9 +107,9 @@ def _chunked_cross_entropy_forward( N_CHUNKS , BLOCK_SIZE : tl.constexpr, DO_SOFTCAPPING : tl.constexpr(tl.int1), - SOFTCAP , + SOFTCAP : tl.constexpr(tl.float32), DO_LOGIT_SCALING : tl.constexpr(tl.int1), - LOGIT_SCALE , + LOGIT_SCALE : tl.constexpr(tl.float32), ): """ 256K vocab divided in 4 chunks @@ -146,14 +146,14 @@ def _chunked_cross_entropy_forward( mask = col_offsets < VOCAB_SIZE label_idx = tl.load(labels_ptr).to(tl.int32) - logits = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")).to(tl.float32) + logits = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")) # Go logit scaling for Cohere: t * x if DO_LOGIT_SCALING: logits = LOGIT_SCALE * logits # Do logit softcapping for Gemma 2: t * tanh(1/t * x) - if DO_SOFTCAPPING: logits = SOFTCAP * triton_tanh(logits / SOFTCAP) + if DO_SOFTCAPPING: logits = SOFTCAP * triton_tanh(logits.to(tl.float32) / SOFTCAP).to(logits.dtype) - # logits = logits.to(tl.float32) + logits = logits.to(tl.float32) c = tl.max(logits, 0) logsumexp = c + tl.log(tl.sum(tl.exp(logits - c), 0)) @@ -175,14 +175,14 @@ 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 , - logits_row_stride , + logits_row_stride : tl.constexpr(tl.int64), dloss_ptr , dloss_row_stride , logsumexp_ptr , @@ -190,9 +190,9 @@ def _cross_entropy_backward( VOCAB_SIZE , BLOCK_SIZE : tl.constexpr, DO_SOFTCAPPING : tl.constexpr(tl.int1), - SOFTCAP , + SOFTCAP : tl.constexpr(tl.float32), DO_LOGIT_SCALING : tl.constexpr(tl.int1), - LOGIT_SCALE , + LOGIT_SCALE : tl.constexpr(tl.float32), ): """ CE_i = -y log(P) = y * (log[sum(exp(x))] - x) @@ -223,7 +223,7 @@ def _cross_entropy_backward( else: dloss = 0.0 - x = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")).to(tl.float32) + x = tl.load(logits_ptr + col_offsets, mask = mask, other = -float("inf")) # Do logit scaling for Cohere if DO_LOGIT_SCALING: @@ -235,12 +235,12 @@ def _cross_entropy_backward( partial = x if DO_SOFTCAPPING: # d/dx [t * tanh(1/t * x)] = 1 - tanh^2(1/t * x) - partial = triton_tanh(x / SOFTCAP) + partial = triton_tanh(x.to(tl.float32) / SOFTCAP).to(x.dtype) x = SOFTCAP * partial pass logsumexp = tl.load(logsumexp_ptr + row_idx) - y = tl.exp(x - logsumexp) + y = tl.exp(x.to(tl.float32) - logsumexp) y = tl.where( col_offsets == label_idx, y - 1.0, # exp(x - logsumexp) - 1 @@ -271,7 +271,7 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): div, mod = divmod(vocab_size, MAX_FUSED_SIZE) n_chunks = div + (mod != 0) - losses = torch.empty(n_rows, dtype = torch.float32, device = logits.device) + 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) @@ -279,7 +279,7 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): if n_chunks == 1: # For small vocabs <= 65336 like Llama, Mistral BLOCK_SIZE, num_warps = calculate_settings(vocab_size) - logsumexp = torch.empty(n_rows, dtype = torch.float32, device = logits.device) + logsumexp = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0") _cross_entropy_forward[(n_rows,)]( logits, logits.stride(0), @@ -296,7 +296,7 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): ) else: # For large vocabs > 65336 like Gemma 256K - logsumexp = torch.empty((n_rows, n_chunks,), dtype = torch.float32, device = logits.device) + logsumexp = torch.empty((n_rows, n_chunks,), dtype = torch.float32, device = "cuda:0") _chunked_cross_entropy_forward[(n_rows, n_chunks,)]( logits, logits.stride(0),