已合并
feat(libdevice): support index_put op #838
candyhong创建于 2025年12月1日
feat(libdevice): support index_put op #838
已合并
candyhong创建于 2025年12月1日
共 8 个文件变更+265-9
@@ -586,6 +586,16 @@ private:
586 static constexpr llvm::StringRef funcNameBase = "triton_embedding_gather";586 static constexpr llvm::StringRef funcNameBase = "triton_embedding_gather";
587};587};
588 588 
589+class IndexPutConverter : public OpConversionPattern<triton::IndexPutOp> {
590+public:
591+ using OpConversionPattern<triton::IndexPutOp>::OpConversionPattern;
592+ LogicalResult
593+ matchAndRewrite(triton::IndexPutOp op, OpAdaptor adaptor,
594+ ConversionPatternRewriter &rewriter) const override;
595+private:
596+ static constexpr llvm::StringRef funcNameBase = "triton_index_put";
597+};
598+ 
589class GatherOutToUbConverter : public OpConversionPattern<triton::GatherOutToUbOp> {599class GatherOutToUbConverter : public OpConversionPattern<triton::GatherOutToUbOp> {
590public:600public:
591 using OpConversionPattern<triton::GatherOutToUbOp>::OpConversionPattern;601 using OpConversionPattern<triton::GatherOutToUbOp>::OpConversionPattern;
@@ -2170,6 +2170,48 @@ EmbeddingGatherConverter::matchAndRewrite(triton::EmbeddingGatherOp op, OpAdapto
2170 return success();2170 return success();
2171}2171}
2172 2172 
2173+LogicalResult
2174+IndexPutConverter::matchAndRewrite(triton::IndexPutOp op, OpAdaptor adaptor,
2175+ ConversionPatternRewriter &rewriter) const
2176+{
2177+ auto loc = op.getLoc();
2178+ 
2179+ auto moduleOp = op->getParentOfType<ModuleOp>();
2180+ rewriter.setInsertionPoint(moduleOp.getBody(),
2181+ std::prev(moduleOp.getBody()->end()));
2182+ 
2183+ auto funcName = generateUniqueFuncName(moduleOp, funcNameBase);
2184+ 
2185+ auto ptr = adaptor.getPtr();
2186+ auto index = op.getIndex();
2187+ auto value = op.getValue();
2188+ auto dim = op.getDim();
2189+ auto dstShape = op.getDstShape();
2190+ auto dstOffset = adaptor.getDstOffset();
2191+ 
2192+ // convert !tt.ptr<f32> to memref<?xf32>
2193+ auto ptrTy = dyn_cast<MemRefType>(ptr.getType());
2194+ if (!ptrTy) {
2195+ return rewriter.notifyMatchFailure(op, "expected MemRefType for ptr");
2196+ }
2197+ SmallVector<Type> inputTypes({ptrTy, index.getType(), value.getType(),
2198+ dim.getType()});
2199+ inputTypes.append(dstShape.getTypes().begin(), dstShape.getTypes().end());
2200+ inputTypes.append(dstOffset.getTypes().begin(), dstOffset.getTypes().end());
2201+ auto libFnType = rewriter.getFunctionType(inputTypes, {});
2202+ auto funcOp = rewriter.create<func::FuncOp>(loc, funcName.str(), libFnType);
2203+ SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private);
2204+ 
2205+ rewriter.setInsertionPoint(op);
2206+ SmallVector<Value> inputVals({ptr, index, value, dim});
2207+ inputVals.append(dstShape.begin(), dstShape.end());
2208+ inputVals.append(dstOffset.begin(), dstOffset.end());
2209+ rewriter.create<func::CallOp>(loc, funcOp.getSymNameAttr(),
2210+ TypeRange({}), inputVals);
2211+ rewriter.eraseOp(op);
2212+ return success();
2213+}
2214+ 
2173LogicalResult2215LogicalResult
2174GatherOutToUbConverter::matchAndRewrite(triton::GatherOutToUbOp op, OpAdaptor adaptor,2216GatherOutToUbConverter::matchAndRewrite(triton::GatherOutToUbOp op, OpAdaptor adaptor,
2175 ConversionPatternRewriter &rewriter) const2217 ConversionPatternRewriter &rewriter) const
@@ -82,6 +82,7 @@ inline bool isSIMTOp(Operation *op)
82{82{
83 return isa<83 return isa<
84 triton::EmbeddingGatherOp,84 triton::EmbeddingGatherOp,
85+ triton::IndexPutOp,
85 triton::GatherOutToUbOp,86 triton::GatherOutToUbOp,
86 triton::IndirectLoadOp,87 triton::IndirectLoadOp,
87 triton::IndirectStoreOp88 triton::IndirectStoreOp
@@ -666,6 +667,7 @@ void TritonToLinalgPass::populateTritonToLinalgConversionPatterns(
666 patterns.add<TTOpConverters::YieldConverter>(patterns.getContext());667 patterns.add<TTOpConverters::YieldConverter>(patterns.getContext());
667 patterns.add<TTOpConverters::GatherConverter>(patterns.getContext());668 patterns.add<TTOpConverters::GatherConverter>(patterns.getContext());
668 patterns.add<TTOpConverters::EmbeddingGatherConverter>(patterns.getContext());669 patterns.add<TTOpConverters::EmbeddingGatherConverter>(patterns.getContext());
670+ patterns.add<TTOpConverters::IndexPutConverter>(patterns.getContext());
669 671 
670 patterns.add<TTOpConverters::DeviceAssertConverter>(patterns.getContext());672 patterns.add<TTOpConverters::DeviceAssertConverter>(patterns.getContext());
671 patterns.add<TTOpConverters::DevicePrintConverter>(patterns.getContext());673 patterns.add<TTOpConverters::DevicePrintConverter>(patterns.getContext());
@@ -736,6 +738,7 @@ void TritonToLinalgPass::annotateTensorKindForModule(ModuleOp moduleOp) {
736 triton::LoadOp>(func);738 triton::LoadOp>(func);
737 // OUTPUT tensors739 // OUTPUT tensors
738 this->walkAndMarkTensorKind<TensorKind::OUTPUT,740 this->walkAndMarkTensorKind<TensorKind::OUTPUT,
741+ triton::IndexPutOp,
739 triton::IndirectStoreOp,742 triton::IndirectStoreOp,
740 triton::StoreOp>(func);743 triton::StoreOp>(func);
741 // INPUT_OUTPUT tensors744 // INPUT_OUTPUT tensors
@@ -392,6 +392,51 @@ def TT_EmbeddingGatherOp : TT_Op<"embedding_gather", [
392 // let hasCanonicalizer = 1;392 // let hasCanonicalizer = 1;
393}393}
394 394 
395+//
396+// IndexPut Op
397+//
398+def TT_IndexPutOp : TT_Op<"index_put", [
399+ MemoryEffects<[MemWrite<GlobalMemory>]>,
400+ SameVariadicOperandSize,
401+]> {
402+ let summary = "Scatter store to a tensor pointer with embedding semantics";
403+ 
404+ let description = [{
405+ Index put values from a tensor into a destination tensor.
406+ 
407+ The operation takes:
408+ - ptr: pointer type, the destination tensor pointer (in GM)
409+ - index: tensor, a index to scatter (in UB)
410+ - value: tensor, a value to store (in UB)
411+ - dim: int, the dimension to scatter along
412+ - dst_shape: tuple of int, the shape of destination tensor
413+ - dst_offset: tuple of int, the offsets of each dimension for destination tensor
414+ 
415+ Constraints:
416+ - `ptr` and `value` must have the same rank.
417+ - `ptr.dtype` only supports `float16`, `bfloat16`, `float32` currently.
418+ - `index` must be an integer tensor, and must be 1D.
419+ - `value` support 2~5D tensors.
420+ - `dim` must be valid (0 <= dim < rank(value) - 1).
421+ }];
422+ 
423+ let arguments = (
424+ ins TT_Ptr:$ptr,
425+ TT_Tensor:$index,
426+ TT_Tensor:$value,
427+ TT_Int:$dim,
428+ Variadic<AnyTypeOf<[I32, I64]>>:$dstShape,
429+ Variadic<AnyTypeOf<[I32, I64]>>:$dstOffset
430+ );
431+ 
432+ let assemblyFormat = [{
433+ $ptr `:` type($ptr) `,` $index `:` type($index) `,`
434+ $value `:` type($value) `,` $dim `:` type($dim) `,`
435+ `[` $dstShape `:` type($dstShape) `]` `,` `[` $dstOffset `:` type($dstOffset) `]`
436+ attr-dict
437+ }];
438+}
439+ 
395//440//
396// GatherOutToUb Op441// GatherOutToUb Op
397//442//
@@ -425,8 +470,8 @@ def TT_GatherOutToUbOp : TT_Op<"gather_out_to_ub", [
425 - `index_tile` must be an integer tensor, with rank between 1 and 5.470 - `index_tile` must be an integer tensor, with rank between 1 and 5.
426 - `dim` must be valid (0 <= dim < rank(index_tile)).471 - `dim` must be valid (0 <= dim < rank(index_tile)).
427 - `other` must be a scalar value.472 - `other` must be a scalar value.
428- - For every dimension `i` not equal to `dim`, `index.size[i]` <= `src.size[i]`.473+ - For every dimension `i` not equal to `dim`, `index_tile.size[i]` <= `src.size[i]`.
429- - The output shape is the same as `index.shape`. If `index` is None, \474+ - The output shape is the same as `index_tile.shape`. If `index_tile` is None, \
430 the output tensor will be an empty tensor with the same shape as `index_tile`.475 the output tensor will be an empty tensor with the same shape as `index_tile`.
431 }];476 }];
432 477 
@@ -1811,6 +1811,16 @@ void init_triton_ir(py::module &&m) {
1811 return self.create<EmbeddingGatherOp>(1811 return self.create<EmbeddingGatherOp>(
1812 resType, src, idx, bound_val, blksiz_val, offsets, numels);1812 resType, src, idx, bound_val, blksiz_val, offsets, numels);
1813 })1813 })
1814+ .def("create_index_put",
1815+ [](TritonOpBuilder &self, Value &ptr, Value &index,
1816+ Value &value, const int32_t dim,
1817+ std::vector<Value> &dstShape, std::vector<Value> &dstOffset) -> void {
1818+ // dim need to be i32 type
1819+ auto dimI32Ty = self.getBuilder().getI32Type();
1820+ auto dim_val = self.create<arith::ConstantIntOp>(dim, dimI32Ty);
1821+ 
1822+ self.create<IndexPutOp>(ptr, index, value, dim_val, dstShape, dstOffset);
1823+ })
1814 .def("create_gather_out_to_ub",1824 .def("create_gather_out_to_ub",
1815 [](TritonOpBuilder &self, Value &src, Value &indexTile, const int64_t indexBoundary,1825 [](TritonOpBuilder &self, Value &src, Value &indexTile, const int64_t indexBoundary,
1816 const int32_t dim, std::vector<Value> &srcStride, std::vector<Value> &indexShape,1826 const int32_t dim, std::vector<Value> &srcStride, std::vector<Value> &indexShape,
@@ -21,7 +21,7 @@
21# THE SOFTWARE.21# THE SOFTWARE.
22 22 
23import os23import os
24-from typing import List, Sequence, Optional, Union24+from typing import List, Sequence, Optional, Union, Tuple
25 25 
26from triton._C.libtriton import ir26from triton._C.libtriton import ir
27from triton.language import semantic as real_semantic27from triton.language import semantic as real_semantic
@@ -645,6 +645,81 @@ def index_select(src: tensor, idx: tensor, bound, lstdim_blksiz, offsets, numels
645 return semantic.embedding_gather(src, idx, bound, lstdim_blksiz, offsets, numels, _builder)645 return semantic.embedding_gather(src, idx, bound, lstdim_blksiz, offsets, numels, _builder)
646 646 
647 647 
648+@builtin
W
Wwangzhanpeng52025年12月2日

之前不是有结论说,非标准的OP,放在libdevice里,也就是 libdevice.index_put 调用。你写在core里,不是变成 tl.index_put 了。

likedislike
candyhong
2025年12月2日 评论:
HaiLijuan
2025年12月2日 评论:
candyhong
2025年12月2日 评论:
KanuaK
KanuaK
2025年12月2日 评论:
649+def index_put(
650+ ptr: tensor,
651+ index: tensor,
652+ value: tensor,
653+ dim: int,
654+ dst_shape: tuple,
655+ dst_offset: tuple,
656+ _builder=None
657+):
658+ """
659+ Index put values from a tensor into a destination tensor.
660+ 
661+ Index put operation for different tensor ranks:
662+ 1. 2D index scatter (dim=0 scatters along rows):
663+ out[index[i]][j] = value[i][j] if dim == 0
664+ out[i][index[j]] = value[i][j] if dim == 1
665+ 2. 3D index scatter (dim=0 scatters along the 0th dimension):
666+ out[index[i]][j][k] = value[i][j][k] if dim == 0
667+ out[i][index[j]][k] = value[i][j][k] if dim == 1
668+ out[i][j][index[k]] = value[i][j][k] if dim == 2
669+ 
670+ :param ptr: pointer type, the destination tensor pointer (in GM)
671+ :param index: tensor, a index to scatter (in UB)
672+ :param value: tensor, a value to store (in UB)
673+ :param dim: int, the dimension to scatter along
674+ :param dst_shape: tuple of int, the shape of destination tensor
675+ :param dst_offset: tuple of int, the offsets of each dimension for destination tensor
676+ 
677+ Constraints
678+ ***********
679+ - `ptr` and `value` must have the same rank.
680+ - `ptr.dtype` only supports `float16`, `bfloat16`, `float32` currently.
681+ - `index` must be an integer tensor, and must be 1D.
682+ - `value` support 2~5D tensors.
683+ - `dim` must be valid (0 <= dim < rank(value) - 1).
684+ 
685+ Example
686+ *******
687+ .. code-block:: python
688+ 
689+ import torch
690+ import triton
691+ import triton.language as tl
692+ from triton.language.extra.ascend.libdevice import index_put
693+ 
694+ @triton.jit
695+ def simple_index_put_kernel(value_ptr, index_ptr, dst_ptr):
696+ # index tile shape: [2]
697+ index_local = tl.arange(0, 2)
698+ x1_local = tl.arange(0, 2)[None, :] # shape=(1,2)
699+ 
700+ index_tile = tl.load(index_ptr + index_local)
701+ value_tile = tl.load(value_ptr + index_local[:, None]*2 + x1_local)
702+ 
703+ index_put(
704+ ptr=dst_ptr,
705+ index=index_tile,
706+ value=value_tile,
707+ dim=0,
708+ dst_shape=(4, 2),
709+ dst_offset=(0, 0)
710+ )
711+ 
712+ dst = torch.zeros((4,2), device='npu', dtype=torch.float32)
713+ value = torch.tensor([[1.,2.], [3.,4.]], device='npu')
714+ index = torch.tensor([2, 0], device='npu')
715+ 
716+ simple_index_put_kernel[(1,)](value, index, dst)
717+ print("IndexPut result:", dst) # ref:[[3.,4.], [0.,0.], [1.,2.], [0.,0.]]
718+ """
719+ dim = _constexpr_to_value(dim)
720+ return semantic.index_put(ptr, index, value, dim, dst_shape, dst_offset, _builder)
721+ 
722+ 
648@builtin723@builtin
649def gather_out_to_ub(724def gather_out_to_ub(
650 src: tensor,725 src: tensor,
@@ -689,8 +764,8 @@ def gather_out_to_ub(
689 - `index_tile` must be an integer tensor, with rank between 1 and 5.764 - `index_tile` must be an integer tensor, with rank between 1 and 5.
690 - `dim` must be valid (0 <= dim < rank(index_tile)).765 - `dim` must be valid (0 <= dim < rank(index_tile)).
691 - `other` must be a scalar value.766 - `other` must be a scalar value.
692- - For every dimension `i` not equal to `dim`, `index.size[i]` <= `src.size[i]`.767+ - For every dimension `i` not equal to `dim`, `index_tile.size[i]` <= `src.size[i]`.
693- - The output shape is the same as `index.shape`. If `index` is None, \768+ - The output shape is the same as `index_tile.shape`. If `index_tile` is None, \
694 the output tensor will be an empty tensor with the same shape as `index_tile`.769 the output tensor will be an empty tensor with the same shape as `index_tile`.
695 770 
696 Example771 Example
@@ -1061,6 +1061,68 @@ def embedding_gather(src: tl.tensor, idx: tl.tensor, bound: int, blksiz: int, of
1061 return wrap_tensor(ret, src.dtype.element_ty, ret_shape)1061 return wrap_tensor(ret, src.dtype.element_ty, ret_shape)
1062 1062 
1063 1063 
1064+def index_put(
1065+ ptr: tl.tensor,
1066+ index: tl.tensor,
1067+ value: tl.tensor,
1068+ dim: int,
1069+ dst_shape: Tuple,
1070+ dst_offset: Tuple,
1071+ builder: ir.builder
1072+):
1073+ """
1074+ Index put values from a tensor into a destination tensor.
1075+ 
1076+ Index put operation for different tensor ranks:
1077+ 1. 2D index scatter (dim=0 scatters along rows):
1078+ out[index[i]][j] = value[i][j] if dim == 0
1079+ out[i][index[j]] = value[i][j] if dim == 1
1080+ 2. 3D index scatter (dim=0 scatters along the 0th dimension):
1081+ out[index[i]][j][k] = value[i][j][k] if dim == 0
1082+ out[i][index[j]][k] = value[i][j][k] if dim == 1
1083+ out[i][j][index[k]] = value[i][j][k] if dim == 2
1084+
1085+ Args:
1086+ - ptr: pointer type, the destination tensor pointer (in GM)
1087+ - index: tensor, a index to scatter (in UB)
1088+ - value: tensor, a value to store (in UB)
1089+ - dim: int, the dimension to scatter along
1090+ - dst_shape: tuple of int, the shape of destination tensor
1091+ - dst_offset: tuple of int, the offsets of each dimension for destination tensor
1092+ 
1093+ Constraints:
1094+ - `ptr` and `value` must have the same rank.
1095+ - `ptr.dtype` only supports `float16`, `bfloat16`, `float32` currently.
1096+ - `index` must be an integer tensor, and must be 1D.
1097+ - `value` support 2~5D tensors.
1098+ - `dim` must be valid (0 <= dim < rank(value) - 1).
1099+ """
1100+ assert index.dtype.is_int(), "index must be an integer tensor"
1101+ if not ptr.dtype.element_ty.is_floating():
1102+ raise ValueError(f"Expected dtype fp16/fp32/bf16, but got {ptr.dtype.element_ty}")
1103+ if not isinstance(dim, int):
1104+ raise ValueError("dim must be of type tl.constexpr")
1105+ 
1106+ v_rank = len(value.shape)
1107+ idx_rank = len(index.shape)
1108+ if idx_rank != 1:
1109+ raise ValueError(f"index rank must be 1, got index rank={idx_rank}")
1110+ if v_rank < 2 or v_rank > 5:
1111+ raise ValueError(f"value rank must be in [2, 5], got value rank={v_rank}")
1112+ if dim < 0 or dim >= v_rank - 1:
1113+ raise ValueError(f"dim must satisfy 0<=dim<value.rank-1 ({v_rank-1}), got dim={dim}")
1114+ 
1115+ require_i64 = index.dtype.is_int64()
1116+ dst_shape = [_convert_elem_to_ir_value(builder, elem, require_i64) for elem in dst_shape]
1117+ dst_offset = [_convert_elem_to_ir_value(builder, elem, require_i64) for elem in dst_offset]
1118+ 
1119+ if len(dst_shape) != v_rank or len(dst_offset) != v_rank:
1120+ raise ValueError(f"len(dst_shape)==len(dst_offset)==value.rank required, "
1121+ f"got {len(dst_shape)}, {len(dst_offset)}, {v_rank}")
1122+ 
1123+ return tl.tensor(builder.create_index_put(ptr.handle, index.handle, value.handle, dim, dst_shape, dst_offset), tl.void)
1124+ 
1125+ 
1064def gather_out_to_ub(1126def gather_out_to_ub(
1065 src: tl.tensor,1127 src: tl.tensor,
1066 index_tile: tl.tensor,1128 index_tile: tl.tensor,
@@ -1104,14 +1166,16 @@ def gather_out_to_ub(
1104 if not src.dtype.element_ty.is_floating():1166 if not src.dtype.element_ty.is_floating():
1105 raise ValueError(f"Expected dtype fp16/fp32/bf16, but got {src.dtype.element_ty}")1167 raise ValueError(f"Expected dtype fp16/fp32/bf16, but got {src.dtype.element_ty}")
1106 1168 
1169+ if not isinstance(index_boundary, int):
1170+ raise ValueError("index_boundary must be of type tl.constexpr")
1171+ if not isinstance(dim, int):
1172+ raise ValueError("dim must be of type tl.constexpr")
1173+ 
1107 idx_rank = len(index_tile.shape)1174 idx_rank = len(index_tile.shape)
1108 if idx_rank < 1 or idx_rank > 5:1175 if idx_rank < 1 or idx_rank > 5:
1109 raise ValueError(f"index_tile rank must be in [1, 5], got rank={idx_rank}")1176 raise ValueError(f"index_tile rank must be in [1, 5], got rank={idx_rank}")
1110 if dim < 0 or dim >= idx_rank:1177 if dim < 0 or dim >= idx_rank:
1111- raise ValueError(f"dim must satisfy 0<=dim<index_tile.rank (index_tile.rank={idx_rank}), got dim={dim}")1178+ raise ValueError(f"dim must satisfy 0<=dim<index_tile.rank ({idx_rank}), got dim={dim}")
1112- if len(src_stride) != idx_rank or len(index_shape) != idx_rank or len(offsets) != idx_rank:
1113- raise ValueError(f"len(src_stride)==len(index_shape)==len(offsets)==index_tile.rank required, "
1114- f"got {len(src_stride)}, {len(index_shape)}, {len(offsets)}, {idx_rank}")
1115 1179 
1116 if other is not None:1180 if other is not None:
1117 other = cast(other, src.dtype.element_ty, _builder)1181 other = cast(other, src.dtype.element_ty, _builder)
@@ -1122,6 +1186,10 @@ def gather_out_to_ub(
1122 index_shape = [_convert_elem_to_ir_value(_builder, elem, False) for elem in index_shape]1186 index_shape = [_convert_elem_to_ir_value(_builder, elem, False) for elem in index_shape]
1123 offsets = [_convert_elem_to_ir_value(_builder, elem, False) for elem in offsets]1187 offsets = [_convert_elem_to_ir_value(_builder, elem, False) for elem in offsets]
1124 1188 
1189+ if len(src_stride) != idx_rank or len(index_shape) != idx_rank or len(offsets) != idx_rank:
1190+ raise ValueError(f"len(src_stride)==len(index_shape)==len(offsets)==index_tile.rank required, "
1191+ f"got {len(src_stride)}, {len(index_shape)}, {len(offsets)}, {idx_rank}")
1192+ 
1125 ret = _builder.create_gather_out_to_ub(1193 ret = _builder.create_gather_out_to_ub(
1126 src.handle,1194 src.handle,
1127 index_tile.handle,1195 index_tile.handle,
@@ -651,6 +651,8 @@ from .triton_patch.language.core import (
651 insert_slice,651 insert_slice,
652 index_select_simd,652 index_select_simd,
653 index_select,653 index_select,
654+ index_put,
655+ gather_out_to_ub,
KanuaK
KanuaKKanuaK2025年12月2日

gather_out_to_ub与L656有重复

likedislike
candyhong
2025年12月2日 评论:
654 gather_out_to_ub,656 gather_out_to_ub,
655 extract_slice,657 extract_slice,
656 trans,658 trans,
@@ -854,4 +856,5 @@ language.extra.ascend.libdevice.fdiv = language.math.fdiv
854language.extra.ascend.libdevice.fma = language.math.fma856language.extra.ascend.libdevice.fma = language.math.fma
855language.extra.ascend.libdevice.abs = language.math.abs857language.extra.ascend.libdevice.abs = language.math.abs
856language.extra.ascend.libdevice.index_select_simd = index_select_simd858language.extra.ascend.libdevice.index_select_simd = index_select_simd
859+language.extra.ascend.libdevice.index_put = index_put
W
Wwangzhanpeng52025年12月2日

没有加ut,能否增加ut测试

likedislike
candyhong
2025年12月2日 评论:
857language.extra.ascend.libdevice.gather_out_to_ub = gather_out_to_ub860language.extra.ascend.libdevice.gather_out_to_ub = gather_out_to_ub