[mlir][vector] Fix bug in extractFromBroadcast folding
extract was incorrectly folded when the source was coming from a broadcast that was both adding new rank and broadcasting the inner dimension. Differential Revision: https://reviews.llvm.org/D123867
This commit is contained in:
parent
64969446bc
commit
b4bcef05b7
@ -1292,20 +1292,25 @@ static Value foldExtractFromBroadcast(ExtractOp extractOp) {
|
||||
};
|
||||
unsigned broadcastSrcRank = getRank(source.getType());
|
||||
unsigned extractResultRank = getRank(extractOp.getType());
|
||||
if (extractResultRank < broadcastSrcRank) {
|
||||
auto extractPos = extractVector<int64_t>(extractOp.getPosition());
|
||||
unsigned rankDiff = broadcastSrcRank - extractResultRank;
|
||||
extractPos.erase(
|
||||
extractPos.begin(),
|
||||
std::next(extractPos.begin(), extractPos.size() - rankDiff));
|
||||
extractOp.setOperand(source);
|
||||
// OpBuilder is only used as a helper to build an I64ArrayAttr.
|
||||
OpBuilder b(extractOp.getContext());
|
||||
extractOp->setAttr(ExtractOp::getPositionAttrStrName(),
|
||||
b.getI64ArrayAttr(extractPos));
|
||||
return extractOp.getResult();
|
||||
}
|
||||
return Value();
|
||||
if (extractResultRank >= broadcastSrcRank)
|
||||
return Value();
|
||||
// Check that the dimension of the result haven't been broadcasted.
|
||||
auto extractVecType = extractOp.getType().dyn_cast<VectorType>();
|
||||
auto broadcastVecType = source.getType().dyn_cast<VectorType>();
|
||||
if (extractVecType && broadcastVecType &&
|
||||
extractVecType.getShape() !=
|
||||
broadcastVecType.getShape().take_back(extractResultRank))
|
||||
return Value();
|
||||
auto extractPos = extractVector<int64_t>(extractOp.getPosition());
|
||||
unsigned rankDiff = broadcastSrcRank - extractResultRank;
|
||||
extractPos.erase(extractPos.begin(),
|
||||
std::next(extractPos.begin(), extractPos.size() - rankDiff));
|
||||
extractOp.setOperand(source);
|
||||
// OpBuilder is only used as a helper to build an I64ArrayAttr.
|
||||
OpBuilder b(extractOp.getContext());
|
||||
extractOp->setAttr(ExtractOp::getPositionAttrStrName(),
|
||||
b.getI64ArrayAttr(extractPos));
|
||||
return extractOp.getResult();
|
||||
}
|
||||
|
||||
// Fold extractOp with source coming from ShapeCast op.
|
||||
|
||||
@ -521,6 +521,17 @@ func @fold_extract_broadcast(%a : f32) -> f32 {
|
||||
|
||||
// -----
|
||||
|
||||
// CHECK-LABEL: fold_extract_broadcast_negative
|
||||
// CHECK: vector.broadcast %{{.*}} : vector<1x1xf32> to vector<1x1x4xf32>
|
||||
// CHECK: vector.extract %{{.*}}[0, 0] : vector<1x1x4xf32>
|
||||
func @fold_extract_broadcast_negative(%a : vector<1x1xf32>) -> vector<4xf32> {
|
||||
%b = vector.broadcast %a : vector<1x1xf32> to vector<1x1x4xf32>
|
||||
%r = vector.extract %b[0, 0] : vector<1x1x4xf32>
|
||||
return %r : vector<4xf32>
|
||||
}
|
||||
|
||||
// -----
|
||||
|
||||
// CHECK-LABEL: fold_extract_splat
|
||||
// CHECK-SAME: %[[A:.*]]: f32
|
||||
// CHECK: return %[[A]] : f32
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user