Skip to content

Commit 1f0ef20

Browse files
authored
Fixed permutation propagation during transpose (#297)
1 parent a7ececb commit 1f0ef20

2 files changed

Lines changed: 37 additions & 12 deletions

File tree

include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp

Lines changed: 19 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -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.
116119
static 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
}

python/examples/conftest.py

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,13 +10,27 @@
1010
def empty_decorator(func):
1111
return func
1212

13+
1314
pytest.mark.interpreter = empty_decorator
1415

1516

1617
@pytest.fixture
1718
def device(request):
1819
return "cpu"
1920

21+
# this fixture is used for test_trans_4d && test_trans_reshape
22+
@pytest.fixture
23+
def with_allocator():
24+
import triton
25+
from triton.runtime._allocation import NullAllocator
26+
from triton._internal_testing import default_alloc_fn
27+
28+
triton.set_allocator(default_alloc_fn)
29+
try:
30+
yield
31+
finally:
32+
triton.set_allocator(NullAllocator())
33+
2034

2135
tests_supported = {
2236
"test_store_eviction_policy",
@@ -56,9 +70,12 @@ def device(request):
5670
"test_load_cache_modifier",
5771
"test_dot_without_load",
5872
"test_cat",
59-
"test_addptr"
73+
"test_addptr",
74+
"test_transpose",
75+
"test_trans_4d",
6076
}
6177

78+
6279
def pytest_collection_modifyitems(config, items):
6380
skip_marker = pytest.mark.skip(reason="CPU backend does not support it yet")
6481
# There is a dependency issue on build machine which breaks bfloat16

0 commit comments

Comments
 (0)