[mlir][tensor] Implement getBufferType for ReshapeOp.
This function should be implemented for ops that work in one-shot bufferization. Reviewed By: springerm Differential Revision: https://reviews.llvm.org/D151548
This commit is contained in:
parent
4b1eb4cf0e
commit
9dbb8eefd4
@ -992,13 +992,28 @@ struct ReshapeOpInterface
|
||||
getBuffer(rewriter, reshapeOp.getShape(), options);
|
||||
if (failed(srcBuffer) || failed(shapeBuffer))
|
||||
return failure();
|
||||
auto resultMemRefType = getMemRefTypeWithStaticIdentityLayout(
|
||||
reshapeOp.getResult().getType(),
|
||||
cast<BaseMemRefType>(srcBuffer->getType()).getMemorySpace());
|
||||
auto maybeResultMemRefType =
|
||||
bufferization::getBufferType(reshapeOp.getResult(), options);
|
||||
if (failed(maybeResultMemRefType))
|
||||
return failure();
|
||||
replaceOpWithNewBufferizedOp<memref::ReshapeOp>(
|
||||
rewriter, op, resultMemRefType, *srcBuffer, *shapeBuffer);
|
||||
rewriter, op, maybeResultMemRefType.value(), *srcBuffer, *shapeBuffer);
|
||||
return success();
|
||||
}
|
||||
|
||||
FailureOr<BaseMemRefType>
|
||||
getBufferType(Operation *op, Value value, const BufferizationOptions &options,
|
||||
const DenseMap<Value, BaseMemRefType> &fixedTypes) const {
|
||||
auto reshapeOp = cast<tensor::ReshapeOp>(op);
|
||||
assert(value == reshapeOp.getResult() && "unexpected value provided");
|
||||
auto maybeSourceBufferType = bufferization::getBufferType(
|
||||
reshapeOp.getSource(), options, fixedTypes);
|
||||
if (failed(maybeSourceBufferType))
|
||||
return failure();
|
||||
return getMemRefTypeWithStaticIdentityLayout(
|
||||
reshapeOp.getResult().getType(),
|
||||
cast<BaseMemRefType>(maybeSourceBufferType.value()).getMemorySpace());
|
||||
}
|
||||
};
|
||||
|
||||
/// Analysis of ParallelInsertSliceOp.
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user