triton-lang/triton

[Question] How to force certain computations to occur in float16?

Open

#1,090 opened on Jan 24, 2023

 (8 comments) (0 reactions) (0 assignees)MLIR (3,126 forks)github user discovery
help wanted

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)?

Contributor guide