Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion src/cuda/tile/_ir/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
23 changes: 22 additions & 1 deletion test/test_gather_scatter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
# ============================================================================
Expand Down