[Question] How to force certain computations to occur in float16?
#1,090 opened on Jan 24, 2023
Repository metrics
- Stars
- (19,982 stars)
- PR merge metrics
- (Avg merge 2d 18h) (185 merged PRs in 30d)
Description
I have the following kernel:
@triton.jit
def unmasked_blend_kernel(
img1_ptr, img2_ptr, ratio: float, total_items: int, BLOCK_SIZE: tl.constexpr
):
block_idx = tl.program_id(0)
offset = block_idx * BLOCK_SIZE
item_element_idxs = tl.arange(0, BLOCK_SIZE)
img1_item_ptrs = img1_ptr + offset + item_element_idxs
img2_item_ptrs = img2_ptr + offset + item_element_idxs
if offset + BLOCK_SIZE > total_items:
mask = item_element_idxs < total_items % BLOCK_SIZE
img1 = tl.load(img1_item_ptrs, mask=mask, other=0)
img2 = tl.load(img2_item_ptrs, mask=mask, other=0)
out = ratio * img1 + (1.0 - ratio) * img2
tl.store(img1_item_ptrs, out, mask=mask)
else:
img1 = tl.load(img1_item_ptrs)
img2 = tl.load(img2_item_ptrs)
out = ratio * img1 + (1.0 - ratio) * img2
tl.store(img1_item_ptrs, out)
Given that this kernel is blending 2 uint8 tensors, I figured it would make sense to do the computation out = ratio * img1 + (1.0 - ratio) * img2 since I don't need high precision.
However, I can't figure out how to force the Triton compiler to do the computations in FP16. I tried doing:
img1 = tl.load(img1_item_ptrs, mask=mask, other=0).to(tl.float16)
img2 = tl.load(img2_item_ptrs, mask=mask, other=0).to(tl.float16)
but from looking at the generated PTX, it just seems like the float16s are just converted to float32s before the multiplication occurs.
Is there a way I can force the Triton compiler to make the multiplications float16? (Or is this impossible for a good reason; i.e maybe this is not what I actually want)?