From 2a3525476e35c62dcd1652be92985b42e9523dfc Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Wed, 23 Sep 2026 18:24:34 -0700 Subject: [PATCH] [Mosaic TPU] Add some fold rules for the transpose op. Reverts 4470f2f1dd91ac973e4270bc47150a97d9644ff6 PiperOrigin-RevId: 987140670 --- .../xla/xla/mosaic/dialect/tpu/tpu_ops.cc | 22 ------------------- .../xla/xla/mosaic/dialect/tpu/tpu_ops.td | 1 - 2 files changed, 23 deletions(-) diff --git a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc index 09432b11e5d..6cf17ae9633 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc +++ b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc @@ -955,28 +955,6 @@ LogicalResult TransposeOp::verify() { return success(); } -OpFoldResult TransposeOp::fold(FoldAdaptor adaptor) { - if (llvm::is_sorted(getPermutation())) { - return getVector(); - } - if (auto cst = dyn_cast_if_present(adaptor.getVector())) { - return cst.reshape(getType()); - } - if (auto prev_op = getVector().getDefiningOp()) { - ArrayRef prev_perm = prev_op.getPermutation(); - ArrayRef curr_perm = getPermutation(); - CHECK_EQ(prev_perm.size(), curr_perm.size()); - SmallVector composed_perm(curr_perm.size()); - for (size_t i = 0; i < curr_perm.size(); ++i) { - composed_perm[i] = prev_perm[curr_perm[i]]; - } - if (llvm::is_sorted(composed_perm)) { - return prev_op.getVector(); - } - } - return nullptr; -} - LogicalResult MemRefBitcastOp::verify() { auto src_ty = getInput().getType(); auto tgt_ty = getType(); diff --git a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td index 1130919108d..e23c2715ba5 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td +++ b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td @@ -1921,7 +1921,6 @@ def TPU_TransposeOp : TPU_Op<"transpose", [Pure]> { } }]; let hasVerifier = 1; - let hasFolder = 1; } def TPU_LogOp : TPU_Op<"log"> {