已合并
feat(op): Add custom op support for __builtin_index_select #1113
candyhong创建于 1月16日
feat(op): Add custom op support for __builtin_index_select #1113
已合并
共 4 个文件变更+132-21
| @@ -0,0 +1,49 @@ | |||
| 1 | +import pytest | ||
| 2 | +import torch | ||
| 3 | +import triton | ||
| 4 | +import triton.language as tl | ||
| 5 | +import triton.language.extra.cann.extension as al | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | +def builtin_index_select_kernel(src_ptr, index_ptr, out_ptr): | ||
W | |||
| 10 | + # Define 2x2 tile indices for output tensor | ||
| 11 | + r = tl.arange(0, 2)[:, None] # Row indices: shape [2, 1] | ||
| 12 | + c = tl.arange(0, 2)[None, :] # Column indices: shape [1, 2] | ||
| 13 | + | ||
| 14 | + # Load index tensor (shape [2]) from GM to UB | ||
| 15 | + idx = tl.load(index_ptr + tl.arange(0, 2)) | ||
| 16 | + # Initialize empty 2x2 output tile in UB (default value: 0) | ||
| 17 | + dst = tl.full((2, 2), 0, dtype=tl.float32) | ||
| 18 | + | ||
| 19 | + # Invoke __builtin_index_select custom op to gather elements | ||
| 20 | + out_tile = al.custom( | ||
| 21 | + "__builtin_index_select", | ||
| 22 | + src_ptr, # Pointer to source tensor in GM | ||
| 23 | + idx, # Index tensor (in UB) for gathering | ||
| 24 | + dim=0, # Dimension to gather along | ||
| 25 | + bound=4, # Upper bound for valid index values (out-of-bound check) | ||
| 26 | + end_offset=(2, 2),# End offsets of each dimension for the index tensor | ||
| 27 | + start_offset=(0, 0), # Start offsets of each dimension for the source tensor | ||
| 28 | + src_stride=(4, 1),# Stride of each dimension for the source tensor in GM | ||
| 29 | + out=dst # Output tensor (in UB) to store gathered elements | ||
| 30 | + ) | ||
| 31 | + | ||
| 32 | + # Store the gathered tile from UB to output tensor in GM | ||
| 33 | + tl.store(out_ptr + r * 2 + c, out_tile) | ||
| 34 | + | ||
| 35 | + | ||
| 36 | +if __name__ == "__main__": | ||
| 37 | + src = torch.tensor( | ||
| 38 | + [[10., 11., 12., 13.], | ||
| 39 | + [20., 21., 22., 23.], | ||
| 40 | + [30., 31., 32., 33.], | ||
| 41 | + [40., 41., 42., 43.]], | ||
| 42 | + device="npu", | ||
| 43 | + dtype=torch.float32, | ||
| 44 | + ) | ||
| 45 | + index = torch.tensor([2, 0], device="npu", dtype=torch.int32) | ||
| 46 | + out = torch.empty((2, 2), device="npu", dtype=torch.float32) | ||
| 47 | + ref = torch.index_select(src, 0, index.to(torch.int64))[:, :2] | ||
| 48 | + builtin_index_select_kernel[(1,)](src, index, out) | ||
| 49 | + torch.testing.assert_close(out, ref) # ref: [[30., 31.], [10., 11.]] | ||
| @@ -15,6 +15,24 @@ def custom_op(builder: ir.builder, op_name: str, **kwargs): | |||
| 15 | raise ValueError(f"Unsupported custom op: {op_name}") | 15 | raise ValueError(f"Unsupported custom op: {op_name}") |
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | +def _is_int_like_elem(x) -> bool: | ||
| 19 | + """Accept int / tl.constexpr(int) / tl.tensor(int*).""" | ||
| 20 | + if isinstance(x, int): | ||
| 21 | + return True | ||
| 22 | + if isinstance(x, tl.constexpr): | ||
| 23 | + # constexpr value should be python int | ||
| 24 | + return isinstance(x.value, int) | ||
| 25 | + if isinstance(x, tl.tensor): | ||
| 26 | + # Offsets/strides must be integer typed (i32/i64 etc.) | ||
| 27 | + return x.dtype.is_int() | ||
| 28 | + return False | ||
| 29 | + | ||
| 30 | + | ||
| 31 | +def _assert_int_like_tuple(name: str, xs): | ||
| 32 | + assert isinstance(xs, (tuple, list)), f"{name} should be a tuple/list, but got {type(xs)}" | ||
| 33 | + assert all(_is_int_like_elem(x) for x in xs), f"{name} should be integer" | ||
| 34 | + | ||
| 35 | + | ||
| 18 | def _convert_elem_to_ir_value(builder, elem, require_i64): | 36 | def _convert_elem_to_ir_value(builder, elem, require_i64): |
| 19 | if isinstance(elem, int): | 37 | if isinstance(elem, int): |
| 20 | elem = tl.constexpr(elem) | 38 | elem = tl.constexpr(elem) |
| @@ -24,41 +24,83 @@ | |||
| 24 | import triton.language.core as tl | 24 | import triton.language.core as tl |
| 25 | from .custom_op import register_custom_op | 25 | from .custom_op import register_custom_op |
| 26 | from .core import CORE, PIPE, MODE | 26 | from .core import CORE, PIPE, MODE |
| 27 | +from ._utils import _is_int_like_elem, _assert_int_like_tuple | ||
| 27 | 28 | ||
| 28 | 29 | ||
| 29 | 30 | ||
| 30 | -class _embedding_gather: | 31 | +class _index_select: |
| 31 | - """This operation take a 2D embedding table in GM and a 1D/2D index tensor in UB, | 32 | + """ |
| 32 | - and produces a 2D/3D output tensor by gathering embedding vectors corresponding | 33 | + This operation gathers values from the src GM tensor into the out UB tensor |
| 33 | - to the index. | 34 | + at positions with offsets specified by the index UB tensor along the specified |
| 35 | + dimension using a SIMT template. This operation supports 2D–5D. | ||
| 34 | 36 | ||
| 35 | Arguments: | 37 | Arguments: |
| 36 | - - src: the embedding table pointer (in GM) | 38 | + - src: pointer type, the source tensor pointer (in GM) |
| 37 | - - index: the index tensor (in UB) | 39 | + - index: tensor, a tensor to gather (in UB) |
| 38 | - - bound: the upper bound of index | 40 | + - dim: int, the dimension to gather along |
| 39 | - - offsets: the offsets of each dimension | 41 | + - bound: int, the upper boundary for index |
| 40 | - - numels: the number of elements of each dimension | 42 | + - end_offset: tuple of int, the end offsets of each dimension for index tensor |
| 41 | - - out: the destination tensor | 43 | + - start_offset: tuple of int, the start offsets of each dimension for src tensor |
| 44 | + - src_stride: tuple of int, the stride of each dimension of src tensor | ||
| 45 | + - other(Optional): scalar value, the default value when index is out of boundary (in UB) | ||
| 46 | + - out: the output tensor (in UB) | ||
| 47 | + | ||
| 48 | + Note: | ||
| 49 | + - Supported source ranks: 2D ~ 5D. | ||
| 50 | + - Supported index ranks: 1D or 2D. | ||
| 51 | + - `dim` must be valid (0 <= dim < source ranks). | ||
| 52 | + | ||
| 53 | + Reference formula: | ||
| 54 | + Index select operation for different tensor ranks: | ||
| 55 | + 1. 2D index gather (0 <= dim <= 1) | ||
| 56 | + 1.1 dim = 0, index_rank = 1, src_rank = 2, out_rank = 2 | ||
| 57 | + index_shape = (Ai,) | ||
| 58 | + end_offset = (Ai_end, B_end) | ||
| 59 | + start_offset = (0, B_begin) | ||
| 60 | + out[i][0:B_end-B_begin] = src[index[i]][B_begin:B_end] | ||
| 61 | + 1.2 dim = 0, index_rank = 2, src_rank = 2, out_rank = 3 | ||
| 62 | + index_shape = (Ai, Aj) | ||
| 63 | + end_offset = (Ai_end, Aj_end, B_end) | ||
| 64 | + start_offset = (0, B_begin) | ||
| 65 | + out[i][j][0:B_end-B_begin] = src[index[i][j]][B_begin:B_end] | ||
| 66 | + 2. 3D index gather (0 <= dim <= 2) | ||
| 67 | + 2.1 dim = 0, index_rank = 2, src_rank = 3, out_rank = 4 | ||
| 68 | + index_shape = (Ai, Aj) | ||
| 69 | + end_offset = (Ai_end, Aj_end, B_end, C_end) | ||
| 70 | + start_offset = (0, B_begin, C_begin) | ||
| 71 | + out[i][j][0:B_end-B_begin][0:C_end-C_begin] = src[index[i][j]][B_begin:B_end][C_begin:C_end] | ||
| 72 | + and so on. | ||
| 42 | """ | 73 | """ |
| 43 | - name = '__builtin_embedding_gather' | 74 | + name = '__builtin_index_select' |
| 44 | core = CORE.VECTOR | 75 | core = CORE.VECTOR |
| 45 | pipe = PIPE.PIPE_V | 76 | pipe = PIPE.PIPE_V |
| 46 | mode = MODE.SIMT | 77 | mode = MODE.SIMT |
| 47 | 78 | ||
| 48 | - def __init__(self, src, index, bound, offsets, numels, out=None): | 79 | + def __init__(self, src, index, dim, bound, end_offset, start_offset, src_stride, other=None, out=None): |
| 49 | assert src.type.is_ptr() or src.dtype.is_ptr(), f"src should be a pointer, but got {src.type}" | 80 | assert src.type.is_ptr() or src.dtype.is_ptr(), f"src should be a pointer, but got {src.type}" |
| 50 | - assert index.dtype.is_int(), "index should be an integer tensor" | 81 | + assert index.dtype.is_int(), "index should be integer tensor" |
| 51 | - assert isinstance(bound, int), "bound should be an integer" | 82 | + src_rank = len(src_stride) |
| 52 | - assert len(offsets) == len(numels), "offsets and numels should have same size" | 83 | + idx_rank = len(index.shape) |
| 53 | - assert all(isinstance(x, int) for x in offsets), "offsets should all be integer" | 84 | + assert 2 <= src_rank <= 5, f"src rank should in [2, 5], but got {src_rank}" |
| 54 | - assert all(isinstance(x, int) for x in numels), "numels should all be integer" | 85 | + assert 1 <= idx_rank <= 2, f"index rank should in [1, 2], but got {idx_rank}" |
| 86 | + assert _is_int_like_elem(dim), "dim should be an integer" | ||
| 87 | + assert _is_int_like_elem(bound), "bound should be an integer" | ||
| 88 | + assert 0 <= dim < src_rank, f"dim should in [0, {src_rank - 1}], but got {dim}" | ||
| 89 | + assert len(start_offset) == len(src_stride), "start_offset and src_stride should have same size" | ||
| 90 | + assert len(end_offset) == idx_rank + len(start_offset) - 1, "len(end_offset) should be equal to index rank + len(start_offset) - 1" | ||
| 91 | + | ||
| 92 | + _assert_int_like_tuple("end_offset", end_offset) | ||
| 93 | + _assert_int_like_tuple("start_offset", start_offset) | ||
| 94 | + _assert_int_like_tuple("src_stride", src_stride) | ||
| 95 | + | ||
| 55 | assert out, "out is required" | 96 | assert out, "out is required" |
| 56 | assert out.dtype == src.dtype.element_ty, "out should have same dtype as src" | 97 | assert out.dtype == src.dtype.element_ty, "out should have same dtype as src" |
| 57 | 98 | ||
| 58 | - # use index type for bound, offsets and numels. | 99 | + # use index type for end_offset, start_offset and src_stride. |
| 59 | - self.arg_type['bound'] = index.dtype | 100 | + self.arg_type['end_offset'] = index.dtype |
| 60 | - self.arg_type['offsets'] = index.dtype | 101 | + self.arg_type['start_offset'] = index.dtype |
| 61 | - self.arg_type['numels'] = index.dtype | 102 | + self.arg_type['src_stride'] = index.dtype |
| 103 | + self.extra_attr = f"src_stride_len={len(src_stride)}" | ||
| 62 | 104 | ||
| 63 | 105 | ||
| 64 | 106 | ||
| @@ -162,6 +162,8 @@ def _make_attrs(op, builder): | |||
| 162 | _add_optional_attr(op, 'symbol', builder, attrs) | 162 | _add_optional_attr(op, 'symbol', builder, attrs) |
| 163 | _add_optional_attr(op, 'source', builder, attrs) | 163 | _add_optional_attr(op, 'source', builder, attrs) |
| 164 | _add_optional_attr(op, 'compile', builder, attrs) | 164 | _add_optional_attr(op, 'compile', builder, attrs) |
| 165 | + # Extra attributes can be added here, such as op.extra_attr="attr_a=xx" | ||
| 166 | + _add_optional_attr(op, 'extra_attr', builder, attrs) | ||
| 165 | return attrs | 167 | return attrs |
| 166 | 168 | ||
| 167 | 169 | ||
这个custom op的 builtin_index_select_kernel 和原来的 get_element/index_select 有什么不同。