已合并
feat(libdevice): support index_put op #838
candyhong创建于 2025年12月1日
feat(libdevice): support index_put op #838
已合并
共 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 | + | ||
| 589 | class GatherOutToUbConverter : public OpConversionPattern<triton::GatherOutToUbOp> { | 599 | class GatherOutToUbConverter : public OpConversionPattern<triton::GatherOutToUbOp> { |
| 590 | public: | 600 | public: |
| 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 | + | ||
| 2173 | LogicalResult | 2215 | LogicalResult |
| 2174 | GatherOutToUbConverter::matchAndRewrite(triton::GatherOutToUbOp op, OpAdaptor adaptor, | 2216 | GatherOutToUbConverter::matchAndRewrite(triton::GatherOutToUbOp op, OpAdaptor adaptor, |
| 2175 | ConversionPatternRewriter &rewriter) const | 2217 | 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::IndirectStoreOp | 88 | 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 tensors | 739 | // 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 tensors | 744 | // 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 Op | 441 | // 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 | ||
| 23 | import os | 23 | import os |
| 24 | -from typing import List, Sequence, Optional, Union | 24 | +from typing import List, Sequence, Optional, Union, Tuple |
| 25 | 25 | ||
| 26 | from triton._C.libtriton import ir | 26 | from triton._C.libtriton import ir |
| 27 | from triton.language import semantic as real_semantic | 27 | from 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 | + | ||
W | |||
| 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 | + | ||
| 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 | 723 | ||
| 649 | def gather_out_to_ub( | 724 | def 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 | Example | 771 | 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 | + | ||
| 1064 | def gather_out_to_ub( | 1126 | def 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, | ||
gather_out_to_ub与L656有重复 ![]() ![]() | |||
| 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 | |||
| 854 | language.extra.ascend.libdevice.fma = language.math.fma | 856 | language.extra.ascend.libdevice.fma = language.math.fma |
| 855 | language.extra.ascend.libdevice.abs = language.math.abs | 857 | language.extra.ascend.libdevice.abs = language.math.abs |
| 856 | language.extra.ascend.libdevice.index_select_simd = index_select_simd | 858 | language.extra.ascend.libdevice.index_select_simd = index_select_simd |
| 859 | +language.extra.ascend.libdevice.index_put = index_put | ||
W 没有加ut,能否增加ut测试 ![]() ![]() | |||
| 857 | language.extra.ascend.libdevice.gather_out_to_ub = gather_out_to_ub | 860 | language.extra.ascend.libdevice.gather_out_to_ub = gather_out_to_ub |


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