Skip to content

Commit 74d4f0e

Browse files
authored
[BugFix] Fix make_int sign-extending negative int8 lanes (#2438)
T.fill of an int8 buffer with a negative scalar (other than -1) silently wrote only 1 of every 4 elements correctly and clobbered the other 3 to -1. The vectorized fill packs four int8 lanes into one 32-bit store via make_int(v,v,v,v); a negative signed char is integer-promoted to int and sign-extended before the shifts/OR, so its high 1-bits flood the neighbouring lanes (e.g. make_int(-7,-7,-7,-7) returned 0xFFFFFFF9 = [-7,-1,-1,-1]). Build the packed value from explicit unsigned bytes so the sign-extension bits cannot leak across lanes. Add an int8 negative-fill regression test. Fixes #2427
1 parent 33994cc commit 74d4f0e

2 files changed

Lines changed: 30 additions & 2 deletions

File tree

‎src/tl_templates/cuda/common.h‎

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -137,10 +137,16 @@ TL_DEVICE unsigned __pack_nv_bfloat162(const bfloat16_t x, const bfloat16_t y) {
137137
return (v1 << 16) | v0;
138138
}
139139

140-
// Pack four char values.
140+
// Pack four char values. Build the 32-bit pattern from unsigned bytes: a
141+
// negative signed char would otherwise sign-extend and flood the other lanes
142+
// through the OR.
141143
TL_DEVICE int make_int(signed char x0, signed char x1, signed char x2,
142144
signed char x3) {
143-
return (x3 << 24) | (x2 << 16) | (x1 << 8) | x0;
145+
const unsigned int b0 = static_cast<unsigned char>(x0);
146+
const unsigned int b1 = static_cast<unsigned char>(x1);
147+
const unsigned int b2 = static_cast<unsigned char>(x2);
148+
const unsigned int b3 = static_cast<unsigned char>(x3);
149+
return static_cast<int>((b3 << 24) | (b2 << 16) | (b1 << 8) | b0);
144150
}
145151

146152
// Pack eight char values.

‎testing/python/language/test_tilelang_language_clear.py‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,5 +81,27 @@ def main(out: T.Tensor((128,), T.float32)):
8181
assert torch.allclose(out, torch.zeros_like(out))
8282

8383

84+
@tilelang.testing.requires_cuda
85+
def test_fill_int8_negative():
86+
M, N = 8, 128
87+
88+
def program(value):
89+
@T.prim_func
90+
def main(out: T.Tensor((M, N), "int8")):
91+
with T.Kernel(1, threads=128):
92+
smem = T.alloc_shared((M, N), "int8")
93+
T.fill(smem, value)
94+
T.copy(smem, out)
95+
96+
return main
97+
98+
for value in (-7, -128, -1, 5):
99+
kernel = tilelang.compile(program(value), out_idx=[0])
100+
out = kernel()
101+
torch.cuda.synchronize()
102+
ref = torch.full((M, N), value, dtype=torch.int8, device="cuda")
103+
torch.testing.assert_close(out, ref)
104+
105+
84106
if __name__ == "__main__":
85107
test_matmul()

0 commit comments

Comments
 (0)