Lines Matching refs:genericOp
372 LogicalResult matchAndRewrite(GenericOp genericOp, in matchAndRewrite() argument
375 for (OpOperand *opOperand : genericOp.getInputAndOutputOperands()) { in matchAndRewrite()
383 rewriter.replaceOp(genericOp, *fusedOpResults); in matchAndRewrite()
453 static bool isFusableWithReshapeByDimExpansion(GenericOp genericOp, in isFusableWithReshapeByDimExpansion() argument
460 return genericOp.hasTensorSemantics() && in isFusableWithReshapeByDimExpansion()
461 llvm::all_of(genericOp.indexing_maps().getValue(), in isFusableWithReshapeByDimExpansion()
467 genericOp.getTiedIndexingMap(fusableOpOperand).getNumResults() > 0 && in isFusableWithReshapeByDimExpansion()
468 llvm::all_of(genericOp.iterator_types(), [](Attribute attr) { in isFusableWithReshapeByDimExpansion()
563 static LogicalResult isGenericOpExpandable(GenericOp genericOp, in isGenericOpExpandable() argument
566 if (!genericOp.hasIndexSemantics()) in isGenericOpExpandable()
575 genericOp, "cannot expand due to index semantics and dynamic dims"); in isGenericOpExpandable()
684 fuseWithReshapeByExpansion(GenericOp genericOp, Operation *reshapeOp, in fuseWithReshapeByExpansion() argument
687 assert(isFusableWithReshapeByDimExpansion(genericOp, fusableOpOperand) && in fuseWithReshapeByExpansion()
702 genericOp, fusableOpOperand, in fuseWithReshapeByExpansion()
708 if (failed(isGenericOpExpandable(genericOp, expansionInfo, rewriter))) in fuseWithReshapeByExpansion()
712 llvm::map_range(genericOp.getIndexingMapsArray(), [&](AffineMap m) { in fuseWithReshapeByExpansion()
717 expandedOpOperands.reserve(genericOp.getNumInputs()); in fuseWithReshapeByExpansion()
718 for (OpOperand *opOperand : genericOp.getInputOperands()) { in fuseWithReshapeByExpansion()
724 if (genericOp.isInputTensor(opOperand)) { in fuseWithReshapeByExpansion()
725 AffineMap indexingMap = genericOp.getTiedIndexingMap(opOperand); in fuseWithReshapeByExpansion()
735 return rewriter.notifyMatchFailure(genericOp, msg); in fuseWithReshapeByExpansion()
742 genericOp.getLoc(), expandedOperandType, opOperand->get(), in fuseWithReshapeByExpansion()
750 Location loc = genericOp.getLoc(); in fuseWithReshapeByExpansion()
752 for (OpOperand *opOperand : genericOp.getOutputOperands()) { in fuseWithReshapeByExpansion()
753 AffineMap indexingMap = genericOp.getTiedIndexingMap(opOperand); in fuseWithReshapeByExpansion()
762 return rewriter.notifyMatchFailure(genericOp, msg); in fuseWithReshapeByExpansion()
769 genericOp.getLoc(), expandedOutputType, opOperand->get(), in fuseWithReshapeByExpansion()
780 rewriter.create<GenericOp>(genericOp.getLoc(), resultTypes, in fuseWithReshapeByExpansion()
784 Region &originalRegion = genericOp->getRegion(0); in fuseWithReshapeByExpansion()
793 for (OpResult opResult : genericOp->getOpResults()) { in fuseWithReshapeByExpansion()
798 genericOp.getTiedIndexingMap( in fuseWithReshapeByExpansion()
799 genericOp.getOutputOperand(resultNumber)), in fuseWithReshapeByExpansion()
802 genericOp.getLoc(), opResult.getType(), in fuseWithReshapeByExpansion()
826 LogicalResult matchAndRewrite(GenericOp genericOp, in matchAndRewrite() argument
828 for (OpOperand *opOperand : genericOp.getInputTensorOperands()) { in matchAndRewrite()
836 if (!isFusableWithReshapeByDimExpansion(genericOp, opOperand) || in matchAndRewrite()
841 fuseWithReshapeByExpansion(genericOp, reshapeOp, opOperand, rewriter); in matchAndRewrite()
844 rewriter.replaceOp(genericOp, *replacementValues); in matchAndRewrite()
1005 getCollapsableIterationSpaceDims(GenericOp genericOp, OpOperand *fusableOperand, in getCollapsableIterationSpaceDims() argument
1008 if (!genericOp.hasTensorSemantics() || genericOp.getNumOutputs() != 1) in getCollapsableIterationSpaceDims()
1011 if (!llvm::all_of(genericOp.getIndexingMapsArray(), [](AffineMap map) { in getCollapsableIterationSpaceDims()
1019 for (const auto &iteratorType : llvm::enumerate(genericOp.iterator_types())) { in getCollapsableIterationSpaceDims()
1026 AffineMap indexingMap = genericOp.getTiedIndexingMap(fusableOperand); in getCollapsableIterationSpaceDims()
1027 auto iteratorTypes = genericOp.iterator_types().getValue(); in getCollapsableIterationSpaceDims()
1088 if (llvm::any_of(genericOp.getIndexingMapsArray(), in getCollapsableIterationSpaceDims()
1271 static Value getCollapsedOpOperand(Location loc, GenericOp genericOp, in getCollapsedOpOperand() argument
1275 AffineMap indexingMap = genericOp.getTiedIndexingMap(opOperand); in getCollapsedOpOperand()
1333 GenericOp genericOp, ArrayRef<ReassociationIndices> foldedIterationDims, in collapseGenericOpIterationDims() argument
1336 if (genericOp.getNumLoops() <= 1 || foldedIterationDims.empty() || in collapseGenericOpIterationDims()
1343 if (failed(collapsingInfo.initialize(genericOp.getNumLoops(), in collapseGenericOpIterationDims()
1346 genericOp, "illegal to collapse specified dimensions"); in collapseGenericOpIterationDims()
1351 genericOp.iterator_types().getValue(), collapsingInfo); in collapseGenericOpIterationDims()
1355 llvm::map_range(genericOp.getIndexingMapsArray(), [&](AffineMap map) { in collapseGenericOpIterationDims()
1359 Location loc = genericOp->getLoc(); in collapseGenericOpIterationDims()
1363 llvm::map_range(genericOp.getInputOperands(), [&](OpOperand *opOperand) { in collapseGenericOpIterationDims()
1364 return getCollapsedOpOperand(loc, genericOp, opOperand, collapsingInfo, in collapseGenericOpIterationDims()
1371 resultTypes.reserve(genericOp.getNumOutputs()); in collapseGenericOpIterationDims()
1372 outputOperands.reserve(genericOp.getNumOutputs()); in collapseGenericOpIterationDims()
1373 for (OpOperand *output : genericOp.getOutputOperands()) { in collapseGenericOpIterationDims()
1375 getCollapsedOpOperand(loc, genericOp, output, collapsingInfo, rewriter); in collapseGenericOpIterationDims()
1384 Block *origOpBlock = &genericOp->getRegion(0).front(); in collapseGenericOpIterationDims()
1394 cast<LinalgOp>(genericOp.getOperation()) in collapseGenericOpIterationDims()
1395 .createLoopRanges(rewriter, genericOp.getLoc()); in collapseGenericOpIterationDims()
1412 for (const auto &originalResult : llvm::enumerate(genericOp->getResults())) { in collapseGenericOpIterationDims()
1420 genericOp.getTiedIndexingMapForResult(originalResult.value()); in collapseGenericOpIterationDims()
1446 LogicalResult matchAndRewrite(GenericOp genericOp, in matchAndRewrite() argument
1448 for (OpOperand *opOperand : genericOp.getInputTensorOperands()) { in matchAndRewrite()
1455 getCollapsableIterationSpaceDims(genericOp, opOperand, in matchAndRewrite()
1463 collapseGenericOpIterationDims(genericOp, collapsableIterationDims, in matchAndRewrite()
1467 genericOp, "failed to do the fusion by collapsing transformation"); in matchAndRewrite()
1470 rewriter.replaceOp(genericOp, *replacements); in matchAndRewrite()
1493 LogicalResult matchAndRewrite(GenericOp genericOp, in matchAndRewrite() argument
1495 if (!genericOp.hasTensorSemantics()) in matchAndRewrite()
1497 for (OpOperand *opOperand : genericOp.getInputOperands()) { in matchAndRewrite()
1536 SmallVector<Location> fusedLocs{genericOp.getLoc()}; in matchAndRewrite()
1537 fusedIndexMaps.reserve(genericOp.getNumInputsAndOutputs()); in matchAndRewrite()
1538 fusedOperands.reserve(genericOp.getNumInputs()); in matchAndRewrite()
1539 fusedLocs.reserve(fusedLocs.size() + genericOp.getNumInputs()); in matchAndRewrite()
1540 for (OpOperand *inputOperand : genericOp.getInputOperands()) { in matchAndRewrite()
1544 fusedIndexMaps.push_back(genericOp.getTiedIndexingMap(inputOperand)); in matchAndRewrite()
1548 for (OpOperand *outputOperand : genericOp.getOutputOperands()) in matchAndRewrite()
1549 fusedIndexMaps.push_back(genericOp.getTiedIndexingMap(outputOperand)); in matchAndRewrite()
1554 genericOp, "fused op loop bound computation failed"); in matchAndRewrite()
1561 SmallVector<Value> outputOperands = genericOp.getOutputOperands(); in matchAndRewrite()
1563 rewriter.getFusedLoc(fusedLocs), genericOp->getResultTypes(), in matchAndRewrite()
1567 genericOp.iterator_types(), in matchAndRewrite()
1573 Region ®ion = genericOp->getRegion(0); in matchAndRewrite()
1581 rewriter.replaceOp(genericOp, fusedOp->getResults()); in matchAndRewrite()
1649 LogicalResult matchAndRewrite(GenericOp genericOp, in matchAndRewrite()
1651 if (!genericOp.hasTensorSemantics()) in matchAndRewrite()
1654 Block &payload = genericOp.region().front(); in matchAndRewrite()
1655 for (OpOperand *opOperand : genericOp.getInputOperands()) { in matchAndRewrite()
1656 if (!genericOp.payloadUsesValueFromOperand(opOperand)) in matchAndRewrite()