triton-lang/triton

Illegal memory access in triton kernel and non-determinism codegen

Open

#1,076 opened on Jan 19, 2023

 (4 comments) (0 reactions) (1 assignee)MLIR (3,127 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

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

Contributor guide