From ce372704ff85a79b3affe22b0f575a1406c31c1f Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 11 Nov 2024 00:17:22 -0800 Subject: [PATCH] triton_cast --- unsloth/kernels/cross_entropy_loss.py | 8 ++++---- unsloth/kernels/utils.py | 6 ++++++ 2 files changed, 10 insertions(+), 4 deletions(-) diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index f82defd405..d347cd1878 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -15,7 +15,7 @@ import triton import triton.language as tl import torch -from .utils import calculate_settings, MAX_FUSED_SIZE, triton_tanh +from .utils import calculate_settings, MAX_FUSED_SIZE, triton_tanh, triton_cast from transformers.models.llama.modeling_llama import logger from packaging.version import Version @@ -64,7 +64,7 @@ def _cross_entropy_forward( This ensures exp(x - max(x))'s maximum is 1 as exp(0) = 1. """ row_idx = tl.program_id(0) - logits_ptr += row_idx * tl.cast(logits_row_stride, tl.int64) + logits_ptr += row_idx * triton_cast(logits_row_stride, tl.int64) loss_ptr += row_idx logsumexp_ptr += row_idx labels_ptr += row_idx @@ -142,7 +142,7 @@ def _chunked_cross_entropy_forward( """ row_idx = tl.program_id(0) chunk_idx = tl.program_id(1) - logits_ptr += row_idx * tl.cast(logits_row_stride, tl.int64) + logits_ptr += row_idx * triton_cast(logits_row_stride, tl.int64) loss_ptr += row_idx logsumexp_ptr += row_idx * N_CHUNKS + chunk_idx labels_ptr += row_idx @@ -216,7 +216,7 @@ def _cross_entropy_backward( row_idx = tl.program_id(0) block_idx = tl.program_id(1) - logits_ptr += row_idx * tl.cast(logits_row_stride, tl.int64) + logits_ptr += row_idx * triton_cast(logits_row_stride, tl.int64) dloss_ptr += row_idx * dloss_row_stride col_offsets = block_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = col_offsets < VOCAB_SIZE diff --git a/unsloth/kernels/utils.py b/unsloth/kernels/utils.py index b394d122fd..cef6ccb864 100644 --- a/unsloth/kernels/utils.py +++ b/unsloth/kernels/utils.py @@ -34,9 +34,15 @@ import triton if Version(triton.__version__) >= Version("3.0.0"): from triton.language.extra import libdevice triton_tanh = libdevice.tanh + triton_cast = tl.cast else: import triton.language as tl triton_tanh = tl.math.tanh + # No casting in old Triton versions + @triton.jit + def triton_cast(x, dtype): + return x.to(dtype) + pass pass