@@ -113,18 +113,29 @@ static SmallVector<utils::IteratorType> getNParallelLoopsAttrs(unsigned n) {
113113 return SmallVector<utils::IteratorType>(n, utils::IteratorType::parallel);
114114}
115115
116+ // if order is empty, transpose the last two dimensions
117+ // otherwise, use the provided order.
118+ // The order must be a permutation of the source rank.
116119static Value getTransposedValue (Value source, const Location loc,
117- ConversionPatternRewriter &rewriter) {
118-
120+ ConversionPatternRewriter &rewriter,
121+ llvm::ArrayRef< int32_t > order = {}) {
119122 auto sourceType = cast<RankedTensorType>(source.getType ());
120123 auto sourceRank = sourceType.getRank ();
121124
122125 SmallVector<int64_t > perm (sourceRank);
123- std::iota (std::begin (perm), std::end (perm), 0 );
124- std::swap (perm[sourceRank - 1 ], perm[sourceRank - 2 ]);
125-
126126 SmallVector<int64_t > transposedShape (sourceType.getShape ());
127- std::swap (transposedShape[sourceRank - 1 ], transposedShape[sourceRank - 2 ]);
127+ if (order.empty ()) {
128+ std::iota (std::begin (perm), std::end (perm), 0 );
129+ std::swap (perm[sourceRank - 1 ], perm[sourceRank - 2 ]);
130+ std::swap (transposedShape[sourceRank - 1 ], transposedShape[sourceRank - 2 ]);
131+ } else {
132+ // Use the provided order
133+ assert (order.size () == sourceRank && " Order size must match source rank" );
134+ for (unsigned i = 0 ; i < sourceRank; ++i) {
135+ perm[i] = order[i];
136+ transposedShape[i] = sourceType.getShape ()[order[i]];
137+ }
138+ }
128139
129140 Value transposeInit = rewriter.create <tensor::EmptyOp>(
130141 loc, transposedShape, sourceType.getElementType ());
@@ -769,11 +780,8 @@ struct TransposeConverter : public OpConversionPattern<triton::TransOp> {
769780 LogicalResult
770781 matchAndRewrite (triton::TransOp op, OpAdaptor adaptor,
771782 ConversionPatternRewriter &rewriter) const override {
772- auto src = adaptor.getSrc ();
773- auto srcRank = cast<ShapedType>(src.getType ()).getRank ();
774- assert (srcRank == 2 && " only expect transposing 2D data" );
775-
776- auto res = getTransposedValue (src, op.getLoc (), rewriter);
783+ auto res = getTransposedValue (adaptor.getSrc (), op.getLoc (), rewriter,
784+ op.getOrder ());
777785 rewriter.replaceOp (op, res);
778786 return success ();
779787 }
0 commit comments