diff --git a/src/cuda/tile/_ir/ops.py b/src/cuda/tile/_ir/ops.py index 4ac86513..fc308cb5 100644 --- a/src/cuda/tile/_ir/ops.py +++ b/src/cuda/tile/_ir/ops.py @@ -1326,7 +1326,8 @@ def generate_bytecode(self, ctx: BytecodeContext) -> tuple[bc.Value, bc.Value]: result_token_type=ctx.type_table.Token, source=ctx.get_value(self.pointer), mask=None if self.mask is None else ctx.get_value(self.mask), - paddingValue=ctx.get_value(self.padding_value), + paddingValue=(None if self.mask is None or self.padding_value is None + else ctx.get_value(self.padding_value)), token=None if self.token is None else ctx.get_value(self.token), memory_ordering_semantics=bc.MemoryOrderingSemantics.WEAK, memory_scope=None, diff --git a/test/test_gather_scatter.py b/test/test_gather_scatter.py index 2f35d4b1..14482d51 100644 --- a/test/test_gather_scatter.py +++ b/test/test_gather_scatter.py @@ -16,7 +16,7 @@ from cuda.tile._ir.cast_ops import _is_implicit_cast_ok from cuda.tile._ir.typing_support import to_dtype from cuda.tile._compile import compile_tile -from util import assert_equal, raises_if +from util import assert_equal, filecheck, raises_if from conftest import float_dtypes, bool_dtypes, int_dtypes, dtype_id from torch.testing import make_tensor @@ -216,6 +216,27 @@ def test_ir_checked_vs_unchecked(kernel, expected_mask): assert (store_ops[0].mask is not None) == expected_mask +@ct.kernel +def load_offset_unmasked(x, y): + ind = ct.arange(8, dtype=ct.int32) + t = x.get_raw_memory().load_offset(ind) + ct.scatter(y, ind, t) + + +@pytest.mark.parametrize("kernel", [copy_8_unchecked, load_offset_unmasked]) +def test_unmasked_load_has_no_padding_operand(kernel): + x = torch.arange(10, 18, dtype=torch.float32, device="cuda:0") + y = torch.zeros_like(x) + sig = ct.compilation.KernelSignature.from_kernel_args( + kernel, (x, y), + ct.compilation.CallingConvention.cutile_python_v1()) + bytecode = compile_tile(kernel._pyfunc, [sig], return_bytecode=True, + return_cubin=False).bytecode + + # The only operand of the load is the pointer: no mask, and so no padding either + filecheck(bytecode, "// CHECK: load_ptr_tko weak %{{[^,]*}} token=") + + # ============================================================================ # Tests for custom mask parameter # ============================================================================