Illegal memory access in triton kernel and non-determinism codegen
#1,076 建立於 2023年1月19日
倉庫指標
- 星標
- (19,995 顆星)
- PR 合併指標
- (平均合併 2天 18小時) (30 天內合併 185 個 PR)
描述
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,
)