已合并
feat(op): Add custom op support for __builtin_index_select #1113
feat(op): Add custom op support for __builtin_index_select #1113
已合并
candyhong创建于 1月16日
共 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+@triton.jit
9+def builtin_index_select_kernel(src_ptr, index_ptr, out_ptr):
W

这个custom op的 builtin_index_select_kernel 和原来的 get_element/index_select 有什么不同。

likedislike
candyhong
1月16日 评论:
candyhong
1月16日 评论:
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+ 
18def _convert_elem_to_ir_value(builder, elem, require_i64):36def _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 @@
24import triton.language.core as tl24import triton.language.core as tl
25from .custom_op import register_custom_op25from .custom_op import register_custom_op
26from .core import CORE, PIPE, MODE26from .core import CORE, PIPE, MODE
27+from ._utils import _is_int_like_elem, _assert_int_like_tuple
27 28 
28 29 
29@register_custom_op30@register_custom_op
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 corresponding33+ 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 index40+ - dim: int, the dimension to gather along
39- - offsets: the offsets of each dimension41+ - bound: int, the upper boundary for index
40- - numels: the number of elements of each dimension42+ - end_offset: tuple of int, the end offsets of each dimension for index tensor
41- - out: the destination tensor43+ - 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.VECTOR75 core = CORE.VECTOR
45 pipe = PIPE.PIPE_V76 pipe = PIPE.PIPE_V
46 mode = MODE.SIMT77 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.dtype100+ self.arg_type['end_offset'] = index.dtype
60- self.arg_type['offsets'] = index.dtype101+ self.arg_type['start_offset'] = index.dtype
61- self.arg_type['numels'] = index.dtype102+ self.arg_type['src_stride'] = index.dtype
103+ self.extra_attr = f"src_stride_len={len(src_stride)}"
62 104 
63 105 
64@register_custom_op106@register_custom_op
@@ -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 attrs167 return attrs
166 168 
167 169