Skip to content

Commit a220e1c

Browse files
[TorchToStablehlo] [TorchToTosa] Add lowering support for torch.aten.triu (#4637)
## Description Add support for lowering `torch.aten.triu` in both `TorchToStablehlo` and `TorchToTosa` conversion passes. To prevent code duplication, the implementation refactors and unifies the core logic of `AtenTrilOp` and `AtenTriuOp` under shared helper templates (`convertTrilOrTriu` and `convertTrilOrTriuTosa`), parameterizing only the comparison direction (GE for triu, LE for tril) and mask generation. ## Tests 1. see unit tests added in the commit. 2. run e2e test via https://gist.github.com/hsqStephenZhang/66e0ed68cc43acdc18a4f4d9959c07a7. all backends generate code as efficient as code generated for `tril` Fixes #4604 --------- Signed-off-by: hsqStephenZhang <stephenzhang666666@gmail.com>
1 parent 3a39627 commit a220e1c

7 files changed

Lines changed: 134 additions & 148 deletions

File tree

lib/Conversion/TorchToStablehlo/Basic.cpp

Lines changed: 31 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -2109,15 +2109,13 @@ LogicalResult ConvertAtenOp<AtenBitwiseRightShiftTensorOp>::matchAndRewrite(
21092109
return success();
21102110
}
21112111

2112-
template <>
2113-
LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewrite(
2114-
AtenTrilOp op, OpAdaptor adaptor,
2115-
ConversionPatternRewriter &rewriter) const {
2116-
2112+
template <typename AtenOpT>
2113+
LogicalResult convertTrilOrTriu(AtenOpT op, Value self, Value diagonal,
2114+
stablehlo::ComparisonDirection dir,
2115+
const TypeConverter *typeConverter,
2116+
ConversionPatternRewriter &rewriter) {
21172117
Location loc = op.getLoc();
21182118

2119-
Value self = adaptor.getSelf();
2120-
21212119
auto selfTy = cast<RankedTensorType>(self.getType());
21222120
if (!selfTy.hasStaticShape()) {
21232121
return op->emitError("dynamic shaped input is not supported");
@@ -2133,25 +2131,23 @@ LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewrite(
21332131
Value rowIdxTensor =
21342132
stablehlo::IotaOp::create(rewriter, loc, iotaTy, 0).getResult();
21352133

2136-
Value diagonal = adaptor.getDiagonal();
21372134
Value diagonalTensor =
21382135
tensor::FromElementsOp::create(rewriter, loc, diagonal).getResult();
21392136

21402137
auto bcastDimensions = rewriter.getDenseI64ArrayAttr({1});
21412138
Value shiftedRowIdxTensor = chlo::BroadcastAddOp::create(
21422139
rewriter, loc, rowIdxTensor, diagonalTensor, bcastDimensions);
21432140

2144-
auto cmpDirectionAttr = stablehlo::ComparisonDirectionAttr::get(
2145-
rewriter.getContext(), stablehlo::ComparisonDirection::LE);
2141+
auto cmpDirectionAttr =
2142+
stablehlo::ComparisonDirectionAttr::get(rewriter.getContext(), dir);
21462143
auto cmpTypeAttr = stablehlo::ComparisonTypeAttr::get(
21472144
rewriter.getContext(), stablehlo::ComparisonType::SIGNED);
21482145
auto cmpTy = iotaTy.clone(rewriter.getI1Type());
21492146
Value cmpRes = stablehlo::CompareOp::create(rewriter, loc, cmpTy,
21502147
colIdxTensor, shiftedRowIdxTensor,
21512148
cmpDirectionAttr, cmpTypeAttr);
21522149

2153-
auto resTy =
2154-
cast<RankedTensorType>(getTypeConverter()->convertType(op.getType()));
2150+
auto resTy = cast<RankedTensorType>(typeConverter->convertType(op.getType()));
21552151

21562152
auto bcastTy = resTy.clone(rewriter.getI1Type());
21572153
auto bcastAttr = rewriter.getDenseI64ArrayAttr({selfRank - 2, selfRank - 1});
@@ -2162,16 +2158,17 @@ LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewrite(
21622158
Value zeroTensor;
21632159
if (isa<mlir::FloatType>(resElemTy)) {
21642160
auto constAttr = SplatElementsAttr::get(
2165-
resTy, llvm::APFloat::getZero(
2166-
cast<FloatType>(resElemTy).getFloatSemantics(), false));
2161+
resTy,
2162+
llvm::APFloat::getZero(
2163+
cast<mlir::FloatType>(resElemTy).getFloatSemantics(), false));
21672164
zeroTensor = stablehlo::ConstantOp::create(rewriter, loc, resTy, constAttr);
21682165
} else if (isa<mlir::IntegerType>(resElemTy)) {
21692166
auto constAttr = SplatElementsAttr::get(
21702167
resTy,
21712168
llvm::APInt::getZero(cast<mlir::IntegerType>(resElemTy).getWidth()));
21722169
zeroTensor = stablehlo::ConstantOp::create(rewriter, loc, resTy, constAttr);
21732170
} else {
2174-
return op.emitError("element type is not float or integer");
2171+
return op->emitError("element type is not float or integer");
21752172
}
21762173

21772174
rewriter.replaceOpWithNewOp<stablehlo::SelectOp>(
@@ -2180,6 +2177,24 @@ LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewrite(
21802177
return success();
21812178
}
21822179

2180+
template <>
2181+
LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewrite(
2182+
AtenTrilOp op, OpAdaptor adaptor,
2183+
ConversionPatternRewriter &rewriter) const {
2184+
return convertTrilOrTriu(op, adaptor.getSelf(), adaptor.getDiagonal(),
2185+
stablehlo::ComparisonDirection::LE,
2186+
getTypeConverter(), rewriter);
2187+
}
2188+
2189+
template <>
2190+
LogicalResult ConvertAtenOp<AtenTriuOp>::matchAndRewrite(
2191+
AtenTriuOp op, OpAdaptor adaptor,
2192+
ConversionPatternRewriter &rewriter) const {
2193+
return convertTrilOrTriu(op, adaptor.getSelf(), adaptor.getDiagonal(),
2194+
stablehlo::ComparisonDirection::GE,
2195+
getTypeConverter(), rewriter);
2196+
}
2197+
21832198
template <>
21842199
LogicalResult ConvertAtenOp<AtenIsfiniteOp>::matchAndRewrite(
21852200
AtenIsfiniteOp op, OpAdaptor adaptor,
@@ -2484,6 +2499,7 @@ void mlir::torch::torch_to_stablehlo::populateBasicOpPatternsAndLegality(
24842499
INSERT_ATENOP_PATTERN(AtenBitwiseRightShiftTensorOp);
24852500

24862501
INSERT_ATENOP_PATTERN(AtenTrilOp);
2502+
INSERT_ATENOP_PATTERN(AtenTriuOp);
24872503
INSERT_ATENOP_PATTERN(AtenIsfiniteOp);
24882504
INSERT_ATENOP_PATTERN(AtenSortOp);
24892505

lib/Conversion/TorchToTosa/TorchToTosa.cpp

Lines changed: 63 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -8795,19 +8795,20 @@ ConvertAtenOp<Aten__InterpolateSizeListScaleListOp>::matchAndRewriteImpl(
87958795
return success();
87968796
}
87978797

8798-
// Template to create supporting tril mask tensor for aten.tril
8798+
// Template to create supporting mask tensor for aten.tril/triu
87998799
template <typename T>
8800-
Value createTrilMask(PatternRewriter &rewriter, Operation *op,
8801-
ArrayRef<int64_t> shape, int64_t h, int64_t w,
8802-
int64_t diagonal) {
8800+
Value createTrilOrTriuMask(PatternRewriter &rewriter, Operation *op,
8801+
ArrayRef<int64_t> shape, int64_t h, int64_t w,
8802+
int64_t diagonal, bool isTril) {
88038803
SmallVector<T> vec;
88048804

88058805
for (int64_t i = 0; i < h; i++) {
88068806
for (int64_t j = 0; j < w; j++) {
88078807
// Positive diagonal value includes as many diagonals above the main
88088808
// diagonal, while negative diagonal value excludes as many diagonals
88098809
// below the main diagonal.
8810-
if (i >= j - diagonal) {
8810+
auto cmp = isTril ? i >= j - diagonal : i <= j - diagonal;
8811+
if (cmp) {
88118812
vec.push_back(static_cast<T>(1));
88128813
} else {
88138814
vec.push_back(static_cast<T>(0));
@@ -8818,13 +8819,10 @@ Value createTrilMask(PatternRewriter &rewriter, Operation *op,
88188819
return tosa::getConstTensor<T>(rewriter, op, vec, shape).value();
88198820
}
88208821

8821-
// Legalization for aten.tril
8822-
template <>
8823-
LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewriteImpl(
8824-
AtenTrilOp op, OpAdaptor adaptor,
8825-
ConversionPatternRewriter &rewriter) const {
8826-
auto self = adaptor.getSelf();
8827-
8822+
template <typename AtenOpT>
8823+
LogicalResult convertTrilOrTriu(AtenOpT op, Value self, Value diagonalVal,
8824+
bool isTril, const TypeConverter *typeConverter,
8825+
ConversionPatternRewriter &rewriter) {
88288826
// Not a ranked tensor type
88298827
auto selfType = dyn_cast<RankedTensorType>(self.getType());
88308828
if (!selfType)
@@ -8841,27 +8839,33 @@ LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewriteImpl(
88418839
return rewriter.notifyMatchFailure(
88428840
op, "Currently only static shapes are supported");
88438841

8844-
const TypeConverter *typeConverter = this->getTypeConverter();
88458842
RankedTensorType resultType = cast<RankedTensorType>(
88468843
typeConverter->convertType(op->getResult(0).getType()));
88478844
if (!resultType)
88488845
return rewriter.notifyMatchFailure(op, "Result type cannot be empty");
88498846

8847+
// clang-format off
88508848
// Get height, width of input tensor, and diagonal arg to create
88518849
// a const mask tensor to multiply with input.
88528850
// This mask tensor has the same height and width of input tensor
8853-
// and consists of 1's for the lower triangle part and 0's for the rest.
8854-
// For example, with h=4, w=6, diagonal=1:
8851+
// and consists of 1's for the lower/higher triangle part and 0's for the rest.
8852+
// For tril with h=4, w=6, diagonal=1:
88558853
// tensor([[1, 1, 0, 0, 0, 0],
88568854
// [1, 1, 1, 0, 0, 0],
88578855
// [1, 1, 1, 1, 0, 0],
88588856
// [1, 1, 1, 1, 1, 0]])
8857+
// For triu with h=4, w=6, diagonal=1:
8858+
// tensor([[0, 1, 1, 1, 1, 1],
8859+
// [0, 0, 1, 1, 1, 1],
8860+
// [0, 0, 0, 1, 1, 1],
8861+
// [0, 0, 0, 0, 1, 1]])
8862+
// clang-format on
88598863
auto selfShape = selfType.getShape();
88608864
int64_t h = selfShape[selfRank - 2];
88618865
int64_t w = selfShape[selfRank - 1];
88628866
int64_t diagonal;
88638867

8864-
if (!matchPattern(op.getDiagonal(), m_TorchConstantInt(&diagonal)))
8868+
if (!matchPattern(diagonalVal, m_TorchConstantInt(&diagonal)))
88658869
return rewriter.notifyMatchFailure(op, "Diagonal value is not an integer");
88668870

88678871
// Define shape for mask tensor based on rank
@@ -8871,39 +8875,56 @@ LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewriteImpl(
88718875
maskShape.push_back(h);
88728876
maskShape.push_back(w);
88738877

8874-
Value trilMask = TypeSwitch<Type, Value>(resultType.getElementType())
8875-
.Case<mlir::FloatType>([&](auto) {
8876-
return createTrilMask<float>(rewriter, op, maskShape,
8877-
h, w, diagonal);
8878-
})
8879-
.Case<mlir::IntegerType>([&](auto intType) {
8880-
switch (intType.getWidth()) {
8881-
case 1:
8882-
return createTrilMask<bool>(rewriter, op, maskShape,
8883-
h, w, diagonal);
8884-
case 32:
8885-
return createTrilMask<int32_t>(
8886-
rewriter, op, maskShape, h, w, diagonal);
8887-
case 64:
8888-
return createTrilMask<int64_t>(
8889-
rewriter, op, maskShape, h, w, diagonal);
8890-
}
8891-
llvm_unreachable("Invalid integer width");
8892-
});
8893-
8894-
if (mlir::tosa::EqualizeRanks(rewriter, op->getLoc(), self, trilMask)
8895-
.failed())
8878+
Value mask =
8879+
TypeSwitch<Type, Value>(resultType.getElementType())
8880+
.Case<mlir::FloatType>([&](auto) {
8881+
return createTrilOrTriuMask<float>(rewriter, op, maskShape, h, w,
8882+
diagonal, isTril);
8883+
})
8884+
.template Case<mlir::IntegerType>([&](auto intType) {
8885+
switch (intType.getWidth()) {
8886+
case 1:
8887+
return createTrilOrTriuMask<bool>(rewriter, op, maskShape, h, w,
8888+
diagonal, isTril);
8889+
case 32:
8890+
return createTrilOrTriuMask<int32_t>(rewriter, op, maskShape, h,
8891+
w, diagonal, isTril);
8892+
case 64:
8893+
return createTrilOrTriuMask<int64_t>(rewriter, op, maskShape, h,
8894+
w, diagonal, isTril);
8895+
}
8896+
llvm_unreachable("Invalid integer width");
8897+
});
8898+
8899+
if (mlir::tosa::EqualizeRanks(rewriter, op->getLoc(), self, mask).failed())
88968900
return rewriter.notifyMatchFailure(
88978901
op, "Failed to equalize ranks among operands and result");
88988902

8899-
auto result =
8900-
tosa::createMulOpAndCast(rewriter, op, resultType, self, trilMask,
8901-
/*shift=*/0);
8903+
auto result = tosa::createMulOpAndCast(rewriter, op, resultType, self, mask,
8904+
/*shift=*/0);
89028905
rewriter.replaceOp(op, result.getResult());
89038906

89048907
return success();
89058908
}
89068909

8910+
// Legalization for aten.triu
8911+
template <>
8912+
LogicalResult ConvertAtenOp<AtenTriuOp>::matchAndRewriteImpl(
8913+
AtenTriuOp op, OpAdaptor adaptor,
8914+
ConversionPatternRewriter &rewriter) const {
8915+
return convertTrilOrTriu(op, adaptor.getSelf(), op.getDiagonal(),
8916+
/*isTril=*/false, getTypeConverter(), rewriter);
8917+
}
8918+
8919+
// Legalization for aten.tril
8920+
template <>
8921+
LogicalResult ConvertAtenOp<AtenTrilOp>::matchAndRewriteImpl(
8922+
AtenTrilOp op, OpAdaptor adaptor,
8923+
ConversionPatternRewriter &rewriter) const {
8924+
return convertTrilOrTriu(op, adaptor.getSelf(), op.getDiagonal(),
8925+
/*isTril=*/true, getTypeConverter(), rewriter);
8926+
}
8927+
89078928
// Legalization for aten.flip
89088929
template <>
89098930
LogicalResult ConvertAtenOp<AtenFlipOp>::matchAndRewriteImpl(
@@ -11756,6 +11777,7 @@ std::set<StringRef> populateTorchToTosaConversionPatternsAndIllegalOps(
1175611777
INSERT_ATENOP_PATTERN(AtenIscloseOp);
1175711778
INSERT_ATENOP_PATTERN(Aten__InterpolateSizeListScaleListOp);
1175811779
INSERT_ATENOP_PATTERN(AtenTrilOp);
11780+
INSERT_ATENOP_PATTERN(AtenTriuOp);
1175911781
INSERT_ATENOP_PATTERN(AtenDiagonalOp);
1176011782
INSERT_ATENOP_PATTERN(AtenIndexSelectOp);
1176111783
INSERT_ATENOP_PATTERN(AtenFlipOp);

lib/Dialect/Torch/Transforms/DecomposeComplexOps.cpp

Lines changed: 0 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -769,75 +769,6 @@ static Value performLastReduceAndPermute(PatternRewriter &rewriter,
769769
return out;
770770
}
771771

772-
namespace {
773-
class DecomposeAtenTriuOp : public OpRewritePattern<AtenTriuOp> {
774-
public:
775-
using OpRewritePattern::OpRewritePattern;
776-
LogicalResult matchAndRewrite(AtenTriuOp op,
777-
PatternRewriter &rewriter) const override {
778-
Location loc = op.getLoc();
779-
Value input = op.getSelf();
780-
auto inputType = cast<BaseTensorType>(input.getType());
781-
if (!inputType.hasSizes() || !inputType.hasDtype()) {
782-
return rewriter.notifyMatchFailure(op, "should have shape and dtype");
783-
}
784-
if (inputType.getSizes().size() < 2) {
785-
return rewriter.notifyMatchFailure(op, "the rank of tensor should >= 2");
786-
}
787-
788-
Value cstZero =
789-
ConstantIntOp::create(rewriter, loc, rewriter.getI64IntegerAttr(0));
790-
Value cstOne =
791-
ConstantIntOp::create(rewriter, loc, rewriter.getI64IntegerAttr(1));
792-
Value none = ConstantNoneOp::create(rewriter, loc);
793-
794-
Value rowSize = getTensorDimSize(rewriter, input, -2);
795-
Value colSize = getTensorDimSize(rewriter, input, -1);
796-
797-
auto si64Type = rewriter.getIntegerType(/*width=*/64, /*isSigned*/ true);
798-
auto int64DtypeInt = getDtypeIntValueForType(rewriter, loc, si64Type);
799-
auto rowArrangeType = getTensorTypeFromShapeValues({rowSize}, si64Type);
800-
auto colArrangeType = getTensorTypeFromShapeValues({colSize}, si64Type);
801-
802-
Value rowArange =
803-
AtenArangeOp::create(rewriter, loc, rowArrangeType, rowSize,
804-
/*dtype=*/int64DtypeInt, /*layout=*/none,
805-
/*device=*/none, /*pin_memory=*/none);
806-
Value colArange =
807-
AtenArangeOp::create(rewriter, loc, colArrangeType, colSize,
808-
/*dtype=*/int64DtypeInt, /*layout=*/none,
809-
/*device=*/none, /*pin_memory=*/none);
810-
811-
auto unsqueezeRowArangeInfo =
812-
unsqueezeTensor(rewriter, op, rowArange, cstOne);
813-
auto unsqueezeColArangeInfo =
814-
unsqueezeTensor(rewriter, op, colArange, cstZero);
815-
816-
if (failed(unsqueezeRowArangeInfo) || failed(unsqueezeColArangeInfo)) {
817-
return rewriter.notifyMatchFailure(op,
818-
"cannot generate unsqueeze tensor");
819-
}
820-
821-
Value unsqueezeRowArange = unsqueezeRowArangeInfo.value();
822-
Value unsqueezeColArange = unsqueezeColArangeInfo.value();
823-
824-
Value unsqueezeRowArangePlusDiagonal =
825-
AtenAddScalarOp::create(rewriter, loc, unsqueezeRowArange.getType(),
826-
unsqueezeRowArange, op.getDiagonal(), cstOne);
827-
828-
auto boolType = rewriter.getI1Type();
829-
auto condType = getTensorTypeFromShapeValues({rowSize, colSize}, boolType);
830-
Value condTensor =
831-
AtenGeTensorOp::create(rewriter, loc, condType, unsqueezeColArange,
832-
unsqueezeRowArangePlusDiagonal);
833-
834-
rewriter.replaceOpWithNewOp<AtenWhereScalarOtherOp>(
835-
op, op.getResult().getType(), condTensor, input, cstZero);
836-
return success();
837-
}
838-
};
839-
} // namespace
840-
841772
/*
842773
This function calculates the number of elements in the lower triangle (below
843774
the main diagonal) of a tensor with dimensions [row, col]. The main diagonal
@@ -13583,7 +13514,6 @@ class DecomposeComplexOpsPass
1358313514
addPatternIfTargetOpIsIllegal<DecomposeAtenTypeAsOp>(patterns);
1358413515
addPatternIfTargetOpIsIllegal<DecomposeAtenTileOp>(patterns);
1358513516
addPatternIfTargetOpIsIllegal<DecomposeAtenReshapeAsOp>(patterns);
13586-
addPatternIfTargetOpIsIllegal<DecomposeAtenTriuOp>(patterns);
1358713517
addPatternIfTargetOpIsIllegal<DecomposeAtenTriuIndicesOp>(patterns);
1358813518
addPatternIfTargetOpIsIllegal<DecomposeAtenTrilIndicesOp>(patterns);
1358913519
addPatternIfTargetOpIsIllegal<DecomposeAtenDeg2radOp>(patterns);

lib/Dialect/Torch/Transforms/LowerToBackendContract.cpp

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -582,7 +582,6 @@ static void markDecomposedOpsAsIllegal(MLIRContext *context,
582582
target.addIllegalOp<AtenTypeAsOp>();
583583
target.addIllegalOp<AtenTileOp>();
584584
target.addIllegalOp<AtenReshapeAsOp>();
585-
target.addIllegalOp<AtenTriuOp>();
586585
target.addIllegalOp<AtenTriuIndicesOp>();
587586
target.addIllegalOp<AtenTrilIndicesOp>();
588587
target.addIllegalOp<AtenDeg2radOp>();

test/Conversion/TorchToStablehlo/basic.mlir

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -360,3 +360,25 @@ func.func @torch.aten.sort(%arg0: !torch.vtensor<[2,3],f32>) -> (!torch.vtensor<
360360
%values, %indices = torch.aten.sort %arg0, %int-1, %true : !torch.vtensor<[2,3],f32>, !torch.int, !torch.bool -> !torch.vtensor<[2,3],f32>, !torch.vtensor<[2,3],si64>
361361
return %values, %indices : !torch.vtensor<[2,3],f32>, !torch.vtensor<[2,3],si64>
362362
}
363+
364+
// -----
365+
366+
// CHECK-LABEL: func.func @torch.aten.triu(
367+
// CHECK-SAME: %[[ARG_0:.*]]: !torch.vtensor<[2,3,5],f32>,
368+
// CHECK-SAME: %[[ARG_1:.*]]: !torch.int) -> !torch.vtensor<[2,3,5],f32>
369+
// CHECK-DAG: %[[VAL_0:.*]] = torch_c.to_builtin_tensor %[[ARG_0]] : !torch.vtensor<[2,3,5],f32> -> tensor<2x3x5xf32>
370+
// CHECK-DAG: %[[VAL_1:.*]] = torch_c.to_i64 %[[ARG_1]]
371+
// CHECK: %[[VAL_2:.*]] = stablehlo.iota dim = 1 : tensor<3x5xi64>
372+
// CHECK: %[[VAL_3:.*]] = stablehlo.iota dim = 0 : tensor<3x5xi64>
373+
// CHECK: %[[VAL_4:.*]] = tensor.from_elements %[[VAL_1]] : tensor<1xi64>
374+
// CHECK: %[[VAL_5:.*]] = chlo.broadcast_add %[[VAL_3]], %[[VAL_4]] {broadcast_dimensions = array<i64: 1>} : (tensor<3x5xi64>, tensor<1xi64>) -> tensor<3x5xi64>
375+
// CHECK: %[[VAL_6:.*]] = stablehlo.compare GE, %[[VAL_2]], %[[VAL_5]], SIGNED : (tensor<3x5xi64>, tensor<3x5xi64>) -> tensor<3x5xi1>
376+
// CHECK: %[[VAL_7:.*]] = stablehlo.broadcast_in_dim %[[VAL_6]], dims = [1, 2] : (tensor<3x5xi1>) -> tensor<2x3x5xi1>
377+
// CHECK: %[[VAL_8:.*]] = stablehlo.constant dense<0.000000e+00> : tensor<2x3x5xf32>
378+
// CHECK: %[[VAL_9:.*]] = stablehlo.select %[[VAL_7]], %[[VAL_0]], %[[VAL_8]] : tensor<2x3x5xi1>, tensor<2x3x5xf32>
379+
// CHECK: %[[VAL_10:.*]] = torch_c.from_builtin_tensor %[[VAL_9]] : tensor<2x3x5xf32> -> !torch.vtensor<[2,3,5],f32>
380+
// CHECK: return %[[VAL_10:.*]] : !torch.vtensor<[2,3,5],f32>
381+
func.func @torch.aten.triu(%arg0: !torch.vtensor<[2,3,5],f32>, %arg1: !torch.int) -> !torch.vtensor<[2,3,5],f32> {
382+
%0 = torch.aten.triu %arg0, %arg1:!torch.vtensor<[2,3,5],f32>, !torch.int -> !torch.vtensor<[2,3,5],f32>
383+
return %0 : !torch.vtensor<[2,3,5],f32>
384+
}

0 commit comments

Comments
 (0)