[sampling] Fix int32 offset overflow in top-k renorm Triton kernels (#35571)
Co-authored-by: Xiaozhu Meng <mxz297@gmail.com>
This commit is contained in:
co-authored by
Xiaozhu Meng
parent
d216737e47
commit
9234e40aed
@@ -25,7 +25,7 @@ def _mask_and_partial_sum_kernel(
|
|||||||
chunk = tl.program_id(1)
|
chunk = tl.program_id(1)
|
||||||
offsets = chunk * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
offsets = chunk * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
mask = offsets < vocab_size
|
mask = offsets < vocab_size
|
||||||
row_offsets = row * vocab_size + offsets
|
row_offsets = row.to(tl.int64) * vocab_size + offsets
|
||||||
|
|
||||||
probs = tl.load(probs_ptr + row_offsets, mask=mask, other=0.0).to(tl.float32)
|
probs = tl.load(probs_ptr + row_offsets, mask=mask, other=0.0).to(tl.float32)
|
||||||
pivot = tl.load(pivots_ptr + row)
|
pivot = tl.load(pivots_ptr + row)
|
||||||
@@ -43,7 +43,7 @@ def _normalize_kernel(
|
|||||||
vocab_size: tl.constexpr,
|
vocab_size: tl.constexpr,
|
||||||
BLOCK_SIZE: tl.constexpr,
|
BLOCK_SIZE: tl.constexpr,
|
||||||
):
|
):
|
||||||
offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
offsets = tl.program_id(0).to(tl.int64) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
mask = offsets < numel
|
mask = offsets < numel
|
||||||
row = offsets // vocab_size
|
row = offsets // vocab_size
|
||||||
values = tl.load(out_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
|
values = tl.load(out_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
|
||||||
|
|||||||
Reference in New Issue
Block a user