better transpose pattern and cleanup
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
NiccoloN
2026-06-03 12:26:31 +02:00
parent 636310d0cb
commit 0a5e73c3ea
8 changed files with 75 additions and 165 deletions
@@ -15,6 +15,10 @@ using namespace mlir;
namespace onnx_mlir {
namespace {
static bool isInsideSpatialComputeRegion(Operation* op) {
return op->getParentOfType<spatial::SpatCompute>() || op->getParentOfType<spatial::SpatComputeBatch>();
}
static Value createTransposeInit(Value input,
RankedTensorType resultType,
ArrayRef<int64_t> permutation,
@@ -102,10 +106,22 @@ struct TransposeToLinalgTranspose : OpConversionPattern<ONNXTransposeOp> {
return success();
}
}
Value init = createTransposeInit(adaptor.getData(), resultType, *permutation, rewriter, transposeOp.getLoc());
Value transposed =
linalg::TransposeOp::create(rewriter, transposeOp.getLoc(), adaptor.getData(), init, *permutation).getResult()[0];
rewriter.replaceOp(transposeOp, transposed);
auto buildTranspose = [&](Value input) -> Value {
Value init = createTransposeInit(input, resultType, *permutation, rewriter, transposeOp.getLoc());
return linalg::TransposeOp::create(rewriter, transposeOp.getLoc(), input, init, *permutation).getResult()[0];
};
if (isInsideSpatialComputeRegion(transposeOp.getOperation())) {
rewriter.replaceOp(transposeOp, buildTranspose(adaptor.getData()));
return success();
}
auto computeOp = createSpatCompute<1>(
rewriter, transposeOp.getLoc(), TypeRange {resultType}, {}, ValueRange {adaptor.getData()}, [&](Value input) {
spatial::SpatYieldOp::create(rewriter, transposeOp.getLoc(), buildTranspose(input));
});
rewriter.replaceOp(transposeOp, computeOp.getResult(0));
return success();
}
};