triton-lang/triton

Illegal memory access in triton kernel and non-determinism codegen

Aberta

#1.076 aberto em 19 de jan. de 2023

 (4 comentários) (0 reação) (1 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

Hi team, we seem to identified an issue that triton is generating non-deterministic codegen results, and in certain cases (small chances) it seems to have illegal memory access due to out of bound access in shared memory. The triton code is fairly simple (shown below). We run the code many times (disabled pytorch's codegen cache) and it seems it produced different shared memory: Most of times: {"name": "triton__0d1d2d3d4", "shared": 2560, "num_warps": 4, "num_stages": 1} Sometimes: {"name": "triton__0d1d2d3d4", "shared": 544, "num_warps": 4, "num_stages": 1}

So sometimes it only asks for 544 bytes for shared memory space which might lead to out of bound access to shared mem. We checked and ttir is the same across the two runs, but llir is different. Wondering if you can shed some lights on this (and whether the codegen is deterministic. Also this is still on llvm IR - not sure if it repros in MLIR.

from torch._inductor.triton_ops.autotune import pointwise

@pointwise(
    size_hints=[262144, 64],
    tile_hint=TileHint.DEFAULT,
    filename="notebook",
    meta={
        "signature": {0: "*bf16", 1: "*fp32", 2: "*bf16", 3: "i32", 4: "i32"},
        "device": 0,
        "constants": {},
        "mutated_arg_names": [],
        "configs": [instance_descriptor(divisible_by_16=(0, 1, 2, 3), equal_to_1=())],
    },
)
@triton.jit
def triton_(
    in_ptr0,
    in_ptr1,
    out_ptr0,
    xnumel,
    ynumel,
    XBLOCK: tl.constexpr,
    YBLOCK: tl.constexpr,
):
    xnumel = 262144
    ynumel = 62
    xoffset = tl.program_id(0) * XBLOCK
    xindex = xoffset + tl.arange(0, XBLOCK)[:, None]
    xmask = xindex < xnumel
    yoffset = tl.program_id(1) * YBLOCK
    yindex = yoffset + tl.arange(0, YBLOCK)[None, :]
    ymask = yindex < ynumel
    x0 = xindex % 512
    x1 = xindex // 512
    y2 = yindex
    x3 = xindex
    tmp0 = tl.load(
        in_ptr0
        + (
            512
            + x0
            + (1536 * x1)
            + (786432 * y2)
            + (786432 * (((x0 + (512 * x1)) // 262144)))
        ),
        xmask & ymask,
    ).to(tl.float32)
    tmp1 = tl.load(in_ptr1 + (512 + x0), xmask)
    tmp2 = tmp1.to(tl.float32)
    tmp3 = tmp0 + tmp2
    tmp4 = 2.8284271247461903
    tmp5 = tmp3 / tmp4
    tl.store(
        out_ptr0 + (y2 + (62 * x3) + tl.zeros([XBLOCK, YBLOCK], tl.int32)),
        tmp5,
        xmask & ymask,
    )

Guia do colaborador