triton_cast
This commit is contained in:
parent
10d4187522
commit
ce372704ff
2 changed files with 10 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue