triton-lang/triton

Reshape from 2D to 3D fails when compiling to Triton IR

Aberta

#641 aberto em 6 de set. de 2022

 (0 comentário) (0 reação) (0 responsável)MLIR (3.136 forks)github user discovery
help wanted

Métricas do repositório

Stars
 (19.995 estrelas)
Métricas de merge de PR
 (Mesclagem média 2d 18h) (185 fundiu PRs em 30d)

Description

Hello, while I was trying to author a Triton kernel that performs part of the Bert attention computation, I ran into the following error:

  File "/home/yidoe/triton/python/triton/code_gen.py", line 1320, in _compile
    name, asm, shared_mem = _triton.code_gen.compile_ttir(backend, generator.module, device, num_warps, num_stages, extern_libs)
IndexError: map::at

My computation is essentially equivalent to this line of code. What I intend to do is to compute a matrix multiplication, reshape the dot product, transpose it, and then store the result to memory.

I have simplified and isolated the issue down to just performing the reshape using the following test code. I have observed that reshaping a 1D tensor from (32,) to (2, 16) is okay; but reshaping a 2D tensor from (4, 32) to (4, 2, 16) triggers the above compile error. Here's how I tested it (for simplicity, it assumes a single program instance):


import torch
import triton
import triton.language as tl


@triton.jit
def reshape3D_kernel(
    a_ptr, b_ptr,
    stride_am, stride_an,
    stride_bm, stride_bp, stride_bq,
    DIM_M: tl.constexpr,
    DIM_N: tl.constexpr,
    DIM_P: tl.constexpr,
    DIM_Q: tl.constexpr,
):
    offs_am = tl.arange(0, DIM_M)
    offs_an = tl.arange(0, DIM_N)
    
    a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_an[None, :] * stride_an)
    
    a = tl.load(a_ptrs)
    aa = tl.reshape(a, (DIM_M, DIM_P, DIM_Q))
    
    offs_bm = tl.arange(0, DIM_M)
    offs_bp = tl.arange(0, DIM_P)
    offs_bq = tl.arange(0, DIM_Q)
    
    b_ptrs = b_ptr + stride_bm * offs_bm[:, None, None] + stride_bp * offs_bp[None, :, None] + stride_bq * offs_bq[None, None, :]
    
    tl.store(b_ptrs, aa)


def reshape3D(a):
    assert a.is_contiguous()
    
    b = torch.empty((DIM_M, DIM_P, DIM_Q), device=a.device, dtype=a.dtype)
    
    reshape3D_kernel[(1,)](
        a, b,
        a.stride(0), a.stride(1),
        b.stride(0), b.stride(1), b.stride(2),
        DIM_M=DIM_M,
        DIM_N=DIM_N,
        DIM_P=DIM_P,
        DIM_Q=DIM_Q,
    )
    
    return b
    
    
@triton.jit
def reshape2D_kernel(
    a_ptr, b_ptr,
    stride_an,
    stride_bp, stride_bq,
    DIM_N: tl.constexpr,
    DIM_P: tl.constexpr,
    DIM_Q: tl.constexpr,
):
    offs_an = tl.arange(0, DIM_N)
    
    a_ptrs = a_ptr + (offs_an[:] * stride_an)
    
    a = tl.load(a_ptrs)
    aa = tl.reshape(a, (DIM_P, DIM_Q))
    
    offs_bp = tl.arange(0, DIM_P)
    offs_bq = tl.arange(0, DIM_Q)
    
    b_ptrs = b_ptr + stride_bp * offs_bp[:, None] + stride_bq * offs_bq[None, :]
    
    tl.store(b_ptrs, aa)
    
    
def reshape2D(a):
    assert a.is_contiguous()
    
    b = torch.empty((DIM_P, DIM_Q), device=a.device, dtype=a.dtype)
    
    reshape2D_kernel[(1,)](
        a, b,
        a.stride(0),
        b.stride(0), b.stride(1),
        DIM_N=DIM_N,
        DIM_P=DIM_P,
        DIM_Q=DIM_Q,
    )
    
    return b
    
    
TEST_3D = True
if TEST_3D:
    DIM_M = 4
    DIM_N = 32
    DIM_P = 2
    DIM_Q = 16
    
    torch.manual_seed(0)
    a = torch.randn((DIM_M, DIM_N), device='cuda', dtype=torch.float16)

    # reshape from (M, N) to (M, P, Q)
    triton_output = reshape3D(a)
    torch_output = a.reshape(DIM_M, DIM_P, DIM_Q)

    print(f"triton_output={triton_output.shape}\n{triton_output}")
    print(f"torch_output={torch_output.shape}\n{torch_output}")
    assert triton.testing.allclose(triton_output, torch_output)
else:
    DIM_N = 32
    DIM_P = 2
    DIM_Q = 16
    
    torch.manual_seed(0)
    a = torch.randn((DIM_N,), device='cuda', dtype=torch.float16)

    # reshape from (N,) to (P, Q)
    triton_output = reshape2D(a)
    torch_output = a.reshape(DIM_P, DIM_Q)

    print(f"triton_output={triton_output.shape}\n{triton_output}")
    print(f"torch_output={torch_output.shape}\n{torch_output}")
    assert triton.testing.allclose(triton_output, torch_output)

This is tested using the latest master branch. Please let me know if it's a bug or if I missed anything. Thanks!

Guia do colaborador