triton_cast

This commit is contained in:
Daniel Han 2024-11-11 00:17:22 -08:00
commit bbe1dda8c0
2 changed files with 10 additions and 4 deletions

View file

@ -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

View file

@ -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