#include "mlir/Dialect/Affine/IR/AffineOps.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/GPU/IR/GPUDialect.h"
#include "mlir/Dialect/GPU/Utils/DistributionUtils.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/Dialect/Vector/Transforms/VectorDistribution.h"
#include "mlir/IR/AffineExpr.h"
#include "mlir/IR/Attributes.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/Interfaces/SideEffectInterfaces.h"
#include "mlir/Transforms/RegionUtils.h"
#include "llvm/ADT/SetVector.h"
#include "llvm/ADT/SmallVectorExtras.h"
#include "llvm/Support/FormatVariadic.h"
#include <utility>
using namespace mlir;
using namespace mlir::vector;
using namespace mlir::gpu;
static AffineMap calculateImplicitMap(VectorType sequentialType,
VectorType distributedType) {
SmallVector<AffineExpr> perm;
perm.reserve(1);
for (unsigned i = 0, e = sequentialType.getRank(); i < e; i++) {
if (sequentialType.getDimSize(i) != distributedType.getDimSize(i))
perm.push_back(getAffineDimExpr(i, distributedType.getContext()));
}
auto map = AffineMap::get(sequentialType.getRank(), 0, perm,
distributedType.getContext());
return map;
}
static int getDistributedDim(VectorType sequentialType,
VectorType distributedType) {
assert(sequentialType.getRank() == distributedType.getRank() &&
"sequential and distributed vector types must have the same rank");
int64_t distributedDim = -1;
for (int64_t i = 0; i < sequentialType.getRank(); ++i) {
if (distributedType.getDimSize(i) != sequentialType.getDimSize(i)) {
assert(distributedDim == -1 && "found multiple distributed dims");
distributedDim = i;
}
}
return distributedDim;
}
namespace {
struct DistributedLoadStoreHelper {
DistributedLoadStoreHelper(Value sequentialVal, Value distributedVal,
Value laneId, Value zero)
: sequentialVal(sequentialVal), distributedVal(distributedVal),
laneId(laneId), zero(zero) {
sequentialVectorType = dyn_cast<VectorType>(sequentialVal.getType());
distributedVectorType = dyn_cast<VectorType>(distributedVal.getType());
if (sequentialVectorType && distributedVectorType)
distributionMap =
calculateImplicitMap(sequentialVectorType, distributedVectorType);
}
Value buildDistributedOffset(RewriterBase &b, Location loc, int64_t index) {
int64_t distributedSize = distributedVectorType.getDimSize(index);
AffineExpr tid = getAffineSymbolExpr(0, b.getContext());
return b.createOrFold<affine::AffineApplyOp>(loc, tid * distributedSize,
ArrayRef<Value>{laneId});
}
Operation *buildStore(RewriterBase &b, Location loc, Value val,
Value buffer) {
assert((val == distributedVal || val == sequentialVal) &&
"Must store either the preregistered distributed or the "
"preregistered sequential value.");
if (!isa<VectorType>(val.getType()))
return memref::StoreOp::create(b, loc, val, buffer, zero);
int64_t rank = sequentialVectorType.getRank();
SmallVector<Value> indices(rank, zero);
if (val == distributedVal) {
for (auto dimExpr : distributionMap.getResults()) {
int64_t index = cast<AffineDimExpr>(dimExpr).getPosition();
indices[index] = buildDistributedOffset(b, loc, index);
}
}
SmallVector<bool> inBounds(indices.size(), true);
return vector::TransferWriteOp::create(
b, loc, val, buffer, indices,
ArrayRef<bool>(inBounds.begin(), inBounds.end()));
}
Value buildLoad(RewriterBase &b, Location loc, Type type, Value buffer) {
if (!isa<VectorType>(type))
return memref::LoadOp::create(b, loc, buffer, zero);
assert((type == distributedVectorType || type == sequentialVectorType) &&
"Must store either the preregistered distributed or the "
"preregistered sequential type.");
SmallVector<Value> indices(sequentialVectorType.getRank(), zero);
if (type == distributedVectorType) {
for (auto dimExpr : distributionMap.getResults()) {
int64_t index = cast<AffineDimExpr>(dimExpr).getPosition();
indices[index] = buildDistributedOffset(b, loc, index);
}
}
SmallVector<bool> inBounds(indices.size(), true);
return vector::TransferReadOp::create(
b, loc, cast<VectorType>(type), buffer, indices,
std::nullopt,
ArrayRef<bool>(inBounds.begin(), inBounds.end()));
}
Value sequentialVal, distributedVal, laneId, zero;
VectorType sequentialVectorType, distributedVectorType;
AffineMap distributionMap;
};
}
static Operation *cloneOpWithOperandsAndTypes(RewriterBase &rewriter,
Location loc, Operation *op,
ArrayRef<Value> operands,
ArrayRef<Type> resultTypes) {
OperationState res(loc, op->getName().getStringRef(), operands, resultTypes,
op->getAttrs());
return rewriter.create(res);
}
namespace {
struct WarpOpToScfIfPattern : public WarpDistributionPattern {
WarpOpToScfIfPattern(MLIRContext *context,
const WarpExecuteOnLane0LoweringOptions &options,
PatternBenefit benefit = 1)
: WarpDistributionPattern(context, benefit), options(options) {}
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
assert(warpOp.getBodyRegion().hasOneBlock() &&
"expected WarpOp with single block");
Block *warpOpBody = &warpOp.getBodyRegion().front();
Location loc = warpOp.getLoc();
OpBuilder::InsertionGuard g(rewriter);
rewriter.setInsertionPoint(warpOp);
Value c0 = arith::ConstantIndexOp::create(rewriter, loc, 0);
Value isLane0 = arith::CmpIOp::create(
rewriter, loc, arith::CmpIPredicate::eq, warpOp.getLaneid(), c0);
auto ifOp = scf::IfOp::create(rewriter, loc, isLane0,
false);
rewriter.eraseOp(ifOp.thenBlock()->getTerminator());
SmallVector<Value> bbArgReplacements;
for (const auto &it : llvm::enumerate(warpOp.getArgs())) {
Value sequentialVal = warpOpBody->getArgument(it.index());
Value distributedVal = it.value();
DistributedLoadStoreHelper helper(sequentialVal, distributedVal,
warpOp.getLaneid(), c0);
rewriter.setInsertionPoint(ifOp);
Value buffer = options.warpAllocationFn(loc, rewriter, warpOp,
sequentialVal.getType());
helper.buildStore(rewriter, loc, distributedVal, buffer);
rewriter.setInsertionPointToStart(ifOp.thenBlock());
bbArgReplacements.push_back(
helper.buildLoad(rewriter, loc, sequentialVal.getType(), buffer));
}
if (!warpOp.getArgs().empty()) {
rewriter.setInsertionPoint(ifOp);
options.warpSyncronizationFn(loc, rewriter, warpOp);
}
rewriter.mergeBlocks(warpOpBody, ifOp.thenBlock(), bbArgReplacements);
SmallVector<Value> replacements;
auto yieldOp = cast<gpu::YieldOp>(ifOp.thenBlock()->getTerminator());
Location yieldLoc = yieldOp.getLoc();
for (const auto &it : llvm::enumerate(yieldOp.getOperands())) {
Value sequentialVal = it.value();
Value distributedVal = warpOp->getResult(it.index());
DistributedLoadStoreHelper helper(sequentialVal, distributedVal,
warpOp.getLaneid(), c0);
rewriter.setInsertionPoint(ifOp);
Value buffer = options.warpAllocationFn(loc, rewriter, warpOp,
sequentialVal.getType());
rewriter.setInsertionPoint(yieldOp);
helper.buildStore(rewriter, loc, sequentialVal, buffer);
rewriter.setInsertionPointAfter(ifOp);
replacements.push_back(
helper.buildLoad(rewriter, loc, distributedVal.getType(), buffer));
}
if (!yieldOp.getOperands().empty()) {
rewriter.setInsertionPointAfter(ifOp);
options.warpSyncronizationFn(loc, rewriter, warpOp);
}
rewriter.eraseOp(yieldOp);
rewriter.setInsertionPointToEnd(ifOp.thenBlock());
scf::YieldOp::create(rewriter, yieldLoc);
rewriter.replaceOp(warpOp, replacements);
return success();
}
private:
const WarpExecuteOnLane0LoweringOptions &options;
};
static VectorType getDistributedType(VectorType originalType, AffineMap map,
int64_t warpSize) {
if (map.getNumResults() == 0)
return originalType;
SmallVector<int64_t> targetShape(originalType.getShape());
for (unsigned i = 0, e = map.getNumResults(); i < e; i++) {
unsigned position = map.getDimPosition(i);
if (targetShape[position] % warpSize != 0) {
if (warpSize % targetShape[position] != 0) {
return VectorType();
}
warpSize /= targetShape[position];
targetShape[position] = 1;
continue;
}
targetShape[position] = targetShape[position] / warpSize;
warpSize = 1;
break;
}
if (warpSize != 1) {
return VectorType();
}
VectorType targetType =
VectorType::get(targetShape, originalType.getElementType());
return targetType;
}
std::tuple<llvm::SmallSetVector<Value, 32>, SmallVector<Type>,
SmallVector<Type>>
getInnerRegionEscapingValues(WarpExecuteOnLane0Op warpOp, Region &innerRegion,
DistributionMapFn distributionMapFn) {
llvm::SmallSetVector<Value, 32> escapingValues;
SmallVector<Type> escapingValueTypes;
SmallVector<Type> escapingValueDistTypes;
if (innerRegion.empty())
return {std::move(escapingValues), std::move(escapingValueTypes),
std::move(escapingValueDistTypes)};
mlir::visitUsedValuesDefinedAbove(innerRegion, [&](OpOperand *operand) {
Operation *parent = operand->get().getParentRegion()->getParentOp();
if (warpOp->isAncestor(parent)) {
if (!escapingValues.insert(operand->get()))
return;
Type distType = operand->get().getType();
if (auto vecType = dyn_cast<VectorType>(distType)) {
AffineMap map = distributionMapFn(operand->get());
distType = getDistributedType(vecType, map, warpOp.getWarpSize());
}
escapingValueTypes.push_back(operand->get().getType());
escapingValueDistTypes.push_back(distType);
}
});
return {std::move(escapingValues), std::move(escapingValueTypes),
std::move(escapingValueDistTypes)};
}
struct WarpOpTransferWrite : public WarpDistributionPattern {
WarpOpTransferWrite(MLIRContext *ctx, DistributionMapFn fn,
unsigned maxNumElementsToExtract, PatternBenefit b = 1)
: WarpDistributionPattern(ctx, b), distributionMapFn(std::move(fn)),
maxNumElementsToExtract(maxNumElementsToExtract) {}
LogicalResult tryDistributeOp(RewriterBase &rewriter,
vector::TransferWriteOp writeOp,
WarpExecuteOnLane0Op warpOp) const {
VectorType writtenVectorType = writeOp.getVectorType();
if (writtenVectorType.getRank() == 0)
return failure();
AffineMap map = distributionMapFn(writeOp.getVector());
VectorType targetType =
getDistributedType(writtenVectorType, map, warpOp.getWarpSize());
if (!targetType)
return failure();
VectorType maskType;
if (writeOp.getMask()) {
if (!writeOp.getPermutationMap().isMinorIdentity())
return failure();
maskType =
getDistributedType(writeOp.getMaskType(), map, warpOp.getWarpSize());
}
vector::TransferWriteOp newWriteOp =
cloneWriteOp(rewriter, warpOp, writeOp, targetType, maskType);
auto newWarpOp =
newWriteOp.getVector().getDefiningOp<WarpExecuteOnLane0Op>();
rewriter.setInsertionPoint(newWriteOp);
SmallVector<OpFoldResult> delinearizedIdSizes;
for (auto [seqSize, distSize] :
llvm::zip_equal(writtenVectorType.getShape(), targetType.getShape())) {
assert(seqSize % distSize == 0 && "Invalid distributed vector shape");
delinearizedIdSizes.push_back(rewriter.getIndexAttr(seqSize / distSize));
}
SmallVector<Value> delinearized;
if (map.getNumResults() > 1) {
delinearized = mlir::affine::AffineDelinearizeIndexOp::create(
rewriter, newWarpOp.getLoc(), newWarpOp.getLaneid(),
delinearizedIdSizes)
.getResults();
} else {
delinearized.append(targetType.getRank(), newWarpOp.getLaneid());
}
AffineMap indexMap = map.compose(newWriteOp.getPermutationMap());
Location loc = newWriteOp.getLoc();
SmallVector<Value> indices(newWriteOp.getIndices().begin(),
newWriteOp.getIndices().end());
for (auto it : llvm::zip(indexMap.getResults(), map.getResults())) {
AffineExpr d0, d1;
bindDims(newWarpOp.getContext(), d0, d1);
auto indexExpr = dyn_cast<AffineDimExpr>(std::get<0>(it));
if (!indexExpr)
continue;
unsigned indexPos = indexExpr.getPosition();
unsigned vectorPos = cast<AffineDimExpr>(std::get<1>(it)).getPosition();
Value laneId = delinearized[vectorPos];
auto scale =
rewriter.getAffineConstantExpr(targetType.getDimSize(vectorPos));
indices[indexPos] = affine::makeComposedAffineApply(
rewriter, loc, d0 + scale * d1, {indices[indexPos], laneId});
}
newWriteOp.getIndicesMutable().assign(indices);
return success();
}
LogicalResult tryExtractOp(RewriterBase &rewriter,
vector::TransferWriteOp writeOp,
WarpExecuteOnLane0Op warpOp) const {
Location loc = writeOp.getLoc();
VectorType vecType = writeOp.getVectorType();
if (vecType.getNumElements() > maxNumElementsToExtract) {
return rewriter.notifyMatchFailure(
warpOp,
llvm::formatv(
"writes more elements ({0}) than allowed to extract ({1})",
vecType.getNumElements(), maxNumElementsToExtract));
}
if (llvm::all_of(warpOp.getOps(),
llvm::IsaPred<vector::TransferWriteOp, gpu::YieldOp>))
return failure();
SmallVector<Value> yieldValues = {writeOp.getVector()};
SmallVector<Type> retTypes = {vecType};
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, yieldValues, retTypes, newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
auto secondWarpOp = WarpExecuteOnLane0Op::create(rewriter, loc, TypeRange(),
newWarpOp.getLaneid(),
newWarpOp.getWarpSize());
Block &body = secondWarpOp.getBodyRegion().front();
rewriter.setInsertionPointToStart(&body);
auto newWriteOp =
cast<vector::TransferWriteOp>(rewriter.clone(*writeOp.getOperation()));
newWriteOp.getValueToStoreMutable().assign(
newWarpOp.getResult(newRetIndices[0]));
rewriter.eraseOp(writeOp);
gpu::YieldOp::create(rewriter, newWarpOp.getLoc());
return success();
}
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
gpu::YieldOp yield = warpOp.getTerminator();
Operation *lastNode = yield->getPrevNode();
auto writeOp = dyn_cast_or_null<vector::TransferWriteOp>(lastNode);
if (!writeOp)
return failure();
Value maybeMask = writeOp.getMask();
if (!llvm::all_of(writeOp->getOperands(), [&](Value value) {
return writeOp.getVector() == value ||
(maybeMask && maybeMask == value) ||
warpOp.isDefinedOutsideOfRegion(value);
}))
return failure();
if (succeeded(tryDistributeOp(rewriter, writeOp, warpOp)))
return success();
if (writeOp.getMask())
return failure();
if (succeeded(tryExtractOp(rewriter, writeOp, warpOp)))
return success();
return failure();
}
private:
vector::TransferWriteOp cloneWriteOp(RewriterBase &rewriter,
WarpExecuteOnLane0Op warpOp,
vector::TransferWriteOp writeOp,
VectorType targetType,
VectorType maybeMaskType) const {
assert(writeOp->getParentOp() == warpOp &&
"write must be nested immediately under warp");
OpBuilder::InsertionGuard g(rewriter);
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp;
if (maybeMaskType) {
newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, ValueRange{writeOp.getVector(), writeOp.getMask()},
TypeRange{targetType, maybeMaskType}, newRetIndices);
} else {
newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, ValueRange{{writeOp.getVector()}},
TypeRange{targetType}, newRetIndices);
}
rewriter.setInsertionPointAfter(newWarpOp);
auto newWriteOp =
cast<vector::TransferWriteOp>(rewriter.clone(*writeOp.getOperation()));
rewriter.eraseOp(writeOp);
newWriteOp.getValueToStoreMutable().assign(
newWarpOp.getResult(newRetIndices[0]));
if (maybeMaskType)
newWriteOp.getMaskMutable().assign(newWarpOp.getResult(newRetIndices[1]));
return newWriteOp;
}
DistributionMapFn distributionMapFn;
unsigned maxNumElementsToExtract = 1;
};
struct WarpOpElementwise : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *yieldOperand = getWarpResult(warpOp, [](Operation *op) {
return OpTrait::hasElementwiseMappableTraits(op);
});
if (!yieldOperand)
return failure();
Operation *elementWise = yieldOperand->get().getDefiningOp();
unsigned operandIndex = yieldOperand->getOperandNumber();
Value distributedVal = warpOp.getResult(operandIndex);
SmallVector<Value> yieldValues;
SmallVector<Type> retTypes;
Location loc = warpOp.getLoc();
for (OpOperand &operand : elementWise->getOpOperands()) {
Type targetType;
if (auto vecType = dyn_cast<VectorType>(distributedVal.getType())) {
auto operandType = cast<VectorType>(operand.get().getType());
targetType =
VectorType::get(vecType.getShape(), operandType.getElementType());
} else {
auto operandType = operand.get().getType();
assert(!isa<VectorType>(operandType) &&
"unexpected yield of vector from op with scalar result type");
targetType = operandType;
}
retTypes.push_back(targetType);
yieldValues.push_back(operand.get());
}
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, yieldValues, retTypes, newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
SmallVector<Value> newOperands(elementWise->getOperands().begin(),
elementWise->getOperands().end());
for (unsigned i : llvm::seq(unsigned(0), elementWise->getNumOperands())) {
newOperands[i] = newWarpOp.getResult(newRetIndices[i]);
}
OpBuilder::InsertionGuard g(rewriter);
rewriter.setInsertionPointAfter(newWarpOp);
Operation *newOp = cloneOpWithOperandsAndTypes(
rewriter, loc, elementWise, newOperands,
{newWarpOp.getResult(operandIndex).getType()});
rewriter.replaceAllUsesWith(newWarpOp.getResult(operandIndex),
newOp->getResult(0));
return success();
}
};
struct WarpOpConstant : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *yieldOperand =
getWarpResult(warpOp, llvm::IsaPred<arith::ConstantOp>);
if (!yieldOperand)
return failure();
auto constantOp = yieldOperand->get().getDefiningOp<arith::ConstantOp>();
auto dense = dyn_cast<SplatElementsAttr>(constantOp.getValue());
if (!dense)
return failure();
rewriter.startOpModification(warpOp);
unsigned operandIndex = yieldOperand->getOperandNumber();
Attribute scalarAttr = dense.getSplatValue<Attribute>();
auto newAttr = DenseElementsAttr::get(
cast<ShapedType>(warpOp.getResult(operandIndex).getType()), scalarAttr);
Location loc = warpOp.getLoc();
rewriter.setInsertionPointAfter(warpOp);
Value distConstant = arith::ConstantOp::create(rewriter, loc, newAttr);
rewriter.replaceAllUsesWith(warpOp.getResult(operandIndex), distConstant);
rewriter.finalizeOpModification(warpOp);
return success();
}
};
struct WarpOpStep final : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *yieldOperand =
getWarpResult(warpOp, llvm::IsaPred<vector::StepOp>);
if (!yieldOperand)
return failure();
const unsigned operandIdx = yieldOperand->getOperandNumber();
auto stepOp = yieldOperand->get().getDefiningOp<vector::StepOp>();
VectorType resTy = stepOp.getResult().getType();
if (resTy.getNumElements() != static_cast<int64_t>(warpOp.getWarpSize()))
return rewriter.notifyMatchFailure(
warpOp,
llvm::formatv("Expected result size ({0}) to be of warp size ({1})",
resTy.getNumElements(), warpOp.getWarpSize()));
VectorType newVecTy =
cast<VectorType>(warpOp.getResult(operandIdx).getType());
rewriter.setInsertionPointAfter(warpOp);
Value laneIdVec = vector::BroadcastOp::create(rewriter, warpOp.getLoc(),
newVecTy, warpOp.getLaneid());
rewriter.replaceAllUsesWith(warpOp.getResult(operandIdx), laneIdVec);
return success();
}
};
struct WarpOpTransferRead : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand = getWarpResult(warpOp, [](Operation *op) {
return isa<vector::TransferReadOp>(op) && op->hasOneUse();
});
if (!operand)
return rewriter.notifyMatchFailure(
warpOp, "warp result is not a vector.transfer_read op");
auto read = operand->get().getDefiningOp<vector::TransferReadOp>();
if (!warpOp.isDefinedOutsideOfRegion(read.getBase()))
return rewriter.notifyMatchFailure(
read, "source must be defined outside of the region");
unsigned operandIndex = operand->getOperandNumber();
Value distributedVal = warpOp.getResult(operandIndex);
SmallVector<Value, 4> indices(read.getIndices().begin(),
read.getIndices().end());
auto sequentialType = cast<VectorType>(read.getResult().getType());
auto distributedType = cast<VectorType>(distributedVal.getType());
AffineMap map = calculateImplicitMap(sequentialType, distributedType);
AffineMap indexMap = map.compose(read.getPermutationMap());
SmallVector<Value> delinearizedIds;
if (!delinearizeLaneId(rewriter, read.getLoc(), sequentialType.getShape(),
distributedType.getShape(), warpOp.getWarpSize(),
warpOp.getLaneid(), delinearizedIds)) {
return rewriter.notifyMatchFailure(
read, "cannot delinearize lane ID for distribution");
}
assert(!delinearizedIds.empty() || map.getNumResults() == 0);
OpBuilder::InsertionGuard g(rewriter);
SmallVector<Value> additionalResults(indices.begin(), indices.end());
SmallVector<Type> additionalResultTypes(indices.size(),
rewriter.getIndexType());
additionalResults.push_back(read.getPadding());
additionalResultTypes.push_back(read.getPadding().getType());
bool hasMask = false;
if (read.getMask()) {
hasMask = true;
if (!mlir::compressUnusedDims(read.getPermutationMap()).isIdentity())
return rewriter.notifyMatchFailure(
read, "non-trivial permutation maps not supported");
VectorType maskType =
getDistributedType(read.getMaskType(), map, warpOp.getWarpSize());
additionalResults.push_back(read.getMask());
additionalResultTypes.push_back(maskType);
}
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, additionalResults, additionalResultTypes,
newRetIndices);
distributedVal = newWarpOp.getResult(operandIndex);
SmallVector<Value> newIndices;
for (int64_t i = 0, e = indices.size(); i < e; ++i)
newIndices.push_back(newWarpOp.getResult(newRetIndices[i]));
rewriter.setInsertionPointAfter(newWarpOp);
for (auto it : llvm::zip_equal(indexMap.getResults(), map.getResults())) {
AffineExpr d0, d1;
bindDims(read.getContext(), d0, d1);
auto indexExpr = dyn_cast<AffineDimExpr>(std::get<0>(it));
if (!indexExpr)
continue;
unsigned indexPos = indexExpr.getPosition();
unsigned vectorPos = cast<AffineDimExpr>(std::get<1>(it)).getPosition();
int64_t scale = distributedType.getDimSize(vectorPos);
newIndices[indexPos] = affine::makeComposedAffineApply(
rewriter, read.getLoc(), d0 + scale * d1,
{newIndices[indexPos], delinearizedIds[vectorPos]});
}
Value newPadding = newWarpOp.getResult(newRetIndices[indices.size()]);
Value newMask =
hasMask ? newWarpOp.getResult(newRetIndices[newRetIndices.size() - 1])
: Value();
auto newRead = vector::TransferReadOp::create(
rewriter, read.getLoc(), distributedVal.getType(), read.getBase(),
newIndices, read.getPermutationMapAttr(), newPadding, newMask,
read.getInBoundsAttr());
rewriter.replaceAllUsesWith(distributedVal, newRead);
return success();
}
};
struct WarpOpDeadResult : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
SmallVector<Type> newResultTypes;
newResultTypes.reserve(warpOp->getNumResults());
SmallVector<Value> newYieldValues;
newYieldValues.reserve(warpOp->getNumResults());
DenseMap<Value, int64_t> dedupYieldOperandPositionMap;
DenseMap<OpResult, int64_t> dedupResultPositionMap;
gpu::YieldOp yield = warpOp.getTerminator();
for (OpResult result : warpOp.getResults()) {
if (result.use_empty())
continue;
Value yieldOperand = yield.getOperand(result.getResultNumber());
auto it = dedupYieldOperandPositionMap.insert(
std::make_pair(yieldOperand, newResultTypes.size()));
dedupResultPositionMap.insert(std::make_pair(result, it.first->second));
if (!it.second)
continue;
newResultTypes.push_back(result.getType());
newYieldValues.push_back(yieldOperand);
}
if (yield.getNumOperands() == newYieldValues.size())
return failure();
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndReplaceReturns(
rewriter, warpOp, newYieldValues, newResultTypes);
newWarpOp.getBody()->walk([&](Operation *op) {
if (isOpTriviallyDead(op))
rewriter.eraseOp(op);
});
SmallVector<Value> newValues;
newValues.reserve(warpOp->getNumResults());
for (OpResult result : warpOp.getResults()) {
if (result.use_empty())
newValues.push_back(Value());
else
newValues.push_back(
newWarpOp.getResult(dedupResultPositionMap.lookup(result)));
}
rewriter.replaceOp(warpOp, newValues);
return success();
}
};
struct WarpOpForwardOperand : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
gpu::YieldOp yield = warpOp.getTerminator();
Value valForwarded;
unsigned resultIndex;
for (OpOperand &operand : yield->getOpOperands()) {
Value result = warpOp.getResult(operand.getOperandNumber());
if (result.use_empty())
continue;
if (!warpOp.getBodyRegion().isAncestor(operand.get().getParentRegion())) {
if (result.getType() != operand.get().getType())
continue;
valForwarded = operand.get();
resultIndex = operand.getOperandNumber();
break;
}
auto arg = dyn_cast<BlockArgument>(operand.get());
if (!arg || arg.getOwner()->getParentOp() != warpOp.getOperation())
continue;
Value warpOperand = warpOp.getArgs()[arg.getArgNumber()];
if (result.getType() != warpOperand.getType())
continue;
valForwarded = warpOperand;
resultIndex = operand.getOperandNumber();
break;
}
if (!valForwarded)
return failure();
rewriter.startOpModification(warpOp);
rewriter.replaceAllUsesWith(warpOp.getResult(resultIndex), valForwarded);
rewriter.finalizeOpModification(warpOp);
return success();
}
};
struct WarpOpBroadcast : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand =
getWarpResult(warpOp, llvm::IsaPred<vector::BroadcastOp>);
if (!operand)
return failure();
unsigned int operandNumber = operand->getOperandNumber();
auto broadcastOp = operand->get().getDefiningOp<vector::BroadcastOp>();
Location loc = broadcastOp.getLoc();
auto destVecType =
cast<VectorType>(warpOp->getResultTypes()[operandNumber]);
Value broadcastSrc = broadcastOp.getSource();
Type broadcastSrcType = broadcastSrc.getType();
if (vector::isBroadcastableTo(broadcastSrcType, destVecType) !=
vector::BroadcastableToResult::Success)
return failure();
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {broadcastSrc}, {broadcastSrcType}, newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value broadcasted = vector::BroadcastOp::create(
rewriter, loc, destVecType, newWarpOp->getResult(newRetIndices[0]));
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber),
broadcasted);
return success();
}
};
struct WarpOpShapeCast : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand =
getWarpResult(warpOp, llvm::IsaPred<vector::ShapeCastOp>);
if (!operand)
return failure();
auto oldCastOp = operand->get().getDefiningOp<vector::ShapeCastOp>();
unsigned int operandNumber = operand->getOperandNumber();
auto castDistributedType =
cast<VectorType>(warpOp->getResultTypes()[operandNumber]);
VectorType castOriginalType = oldCastOp.getSourceVectorType();
VectorType castResultType = castDistributedType;
unsigned castDistributedRank = castDistributedType.getRank();
unsigned castOriginalRank = castOriginalType.getRank();
if (castDistributedRank < castOriginalRank) {
SmallVector<int64_t> shape(castOriginalRank - castDistributedRank, 1);
llvm::append_range(shape, castDistributedType.getShape());
castDistributedType =
VectorType::get(shape, castDistributedType.getElementType());
}
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {oldCastOp.getSource()}, {castDistributedType},
newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value newCast = vector::ShapeCastOp::create(
rewriter, oldCastOp.getLoc(), castResultType,
newWarpOp->getResult(newRetIndices[0]));
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber), newCast);
return success();
}
};
struct WarpOpCreateMask : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *yieldOperand =
getWarpResult(warpOp, llvm::IsaPred<vector::CreateMaskOp>);
if (!yieldOperand)
return failure();
auto mask = yieldOperand->get().getDefiningOp<vector::CreateMaskOp>();
if (!llvm::all_of(mask->getOperands(), [&](Value value) {
return warpOp.isDefinedOutsideOfRegion(value);
}))
return failure();
Location loc = mask.getLoc();
unsigned operandIndex = yieldOperand->getOperandNumber();
auto distType = cast<VectorType>(warpOp.getResult(operandIndex).getType());
VectorType seqType = mask.getVectorType();
ArrayRef<int64_t> seqShape = seqType.getShape();
ArrayRef<int64_t> distShape = distType.getShape();
rewriter.setInsertionPointAfter(warpOp);
SmallVector<Value> delinearizedIds;
if (!delinearizeLaneId(rewriter, loc, seqShape, distShape,
warpOp.getWarpSize(), warpOp.getLaneid(),
delinearizedIds))
return rewriter.notifyMatchFailure(
mask, "cannot delinearize lane ID for distribution");
assert(!delinearizedIds.empty());
rewriter.startOpModification(warpOp);
AffineExpr s0, s1;
bindSymbols(rewriter.getContext(), s0, s1);
SmallVector<Value> newOperands;
for (int i = 0, e = distShape.size(); i < e; ++i) {
Value maskDimIdx = affine::makeComposedAffineApply(
rewriter, loc, s1 - s0 * distShape[i],
{delinearizedIds[i], mask.getOperand(i)});
newOperands.push_back(maskDimIdx);
}
auto newMask =
vector::CreateMaskOp::create(rewriter, loc, distType, newOperands);
rewriter.replaceAllUsesWith(warpOp.getResult(operandIndex), newMask);
rewriter.finalizeOpModification(warpOp);
return success();
}
};
struct WarpOpInsertStridedSlice : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand =
getWarpResult(warpOp, llvm::IsaPred<vector::InsertStridedSliceOp>);
if (!operand)
return failure();
unsigned int operandNumber = operand->getOperandNumber();
auto insertOp =
operand->get().getDefiningOp<vector::InsertStridedSliceOp>();
auto distributedType =
cast<VectorType>(warpOp.getResult(operandNumber).getType());
if (distributedType.getRank() < 2)
return rewriter.notifyMatchFailure(
insertOp, "result vector type must be 2D or higher");
auto yieldedType = cast<VectorType>(operand->get().getType());
int64_t destDistributedDim =
getDistributedDim(yieldedType, distributedType);
assert(destDistributedDim != -1 && "could not find distributed dimension");
VectorType srcType = insertOp.getSourceVectorType();
VectorType destType = insertOp.getDestVectorType();
int64_t sourceDistributedDim =
destDistributedDim - (destType.getRank() - srcType.getRank());
if (sourceDistributedDim < 0)
return rewriter.notifyMatchFailure(
insertOp,
"distributed dimension must be in the last k dims of dest vector");
if (srcType.getDimSize(sourceDistributedDim) !=
destType.getDimSize(destDistributedDim))
return rewriter.notifyMatchFailure(
insertOp, "distributed dimension must be fully inserted");
SmallVector<int64_t> newSourceDistShape(
insertOp.getSourceVectorType().getShape());
newSourceDistShape[sourceDistributedDim] =
distributedType.getDimSize(destDistributedDim);
auto newSourceTy =
VectorType::get(newSourceDistShape, distributedType.getElementType());
VectorType newDestTy = distributedType;
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {insertOp.getValueToStore(), insertOp.getDest()},
{newSourceTy, newDestTy}, newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value distributedSource = newWarpOp->getResult(newRetIndices[0]);
Value distributedDest = newWarpOp->getResult(newRetIndices[1]);
Value newInsert = vector::InsertStridedSliceOp::create(
rewriter, insertOp.getLoc(), distributedDest.getType(),
distributedSource, distributedDest, insertOp.getOffsets(),
insertOp.getStrides());
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber), newInsert);
return success();
}
};
struct WarpOpExtractStridedSlice : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand =
getWarpResult(warpOp, llvm::IsaPred<vector::ExtractStridedSliceOp>);
if (!operand)
return failure();
unsigned int operandNumber = operand->getOperandNumber();
auto extractOp =
operand->get().getDefiningOp<vector::ExtractStridedSliceOp>();
auto distributedType =
cast<VectorType>(warpOp.getResult(operandNumber).getType());
if (distributedType.getRank() < 2)
return rewriter.notifyMatchFailure(
extractOp, "result vector type must be 2D or higher");
auto yieldedType = cast<VectorType>(operand->get().getType());
int64_t distributedDim = getDistributedDim(yieldedType, distributedType);
assert(distributedDim != -1 && "could not find distributed dimension");
int64_t numOfExtractedDims =
static_cast<int64_t>(extractOp.getSizes().size());
if (distributedDim < numOfExtractedDims) {
int64_t distributedDimOffset =
llvm::cast<IntegerAttr>(extractOp.getOffsets()[distributedDim])
.getInt();
int64_t distributedDimSize =
llvm::cast<IntegerAttr>(extractOp.getSizes()[distributedDim])
.getInt();
if (distributedDimOffset != 0 ||
distributedDimSize != yieldedType.getDimSize(distributedDim))
return rewriter.notifyMatchFailure(
extractOp, "distributed dimension must be fully extracted");
}
SmallVector<int64_t> newDistributedShape(
extractOp.getSourceVectorType().getShape());
newDistributedShape[distributedDim] =
distributedType.getDimSize(distributedDim);
auto newDistributedType =
VectorType::get(newDistributedShape, distributedType.getElementType());
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {extractOp.getSource()}, {newDistributedType},
newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
SmallVector<Attribute> distributedSizes = llvm::map_to_vector(
extractOp.getSizes(), [](Attribute attr) { return attr; });
if (distributedDim < static_cast<int64_t>(distributedSizes.size()))
distributedSizes[distributedDim] = rewriter.getI64IntegerAttr(
distributedType.getDimSize(distributedDim));
Value distributedVec = newWarpOp->getResult(newRetIndices[0]);
Value newExtract = vector::ExtractStridedSliceOp::create(
rewriter, extractOp.getLoc(), distributedType, distributedVec,
extractOp.getOffsets(),
ArrayAttr::get(rewriter.getContext(), distributedSizes),
extractOp.getStrides());
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber),
newExtract);
return success();
}
};
struct WarpOpExtract : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand =
getWarpResult(warpOp, llvm::IsaPred<vector::ExtractOp>);
if (!operand)
return failure();
unsigned int operandNumber = operand->getOperandNumber();
auto extractOp = operand->get().getDefiningOp<vector::ExtractOp>();
VectorType extractSrcType = extractOp.getSourceVectorType();
Location loc = extractOp.getLoc();
if (extractSrcType.getRank() <= 1) {
return failure();
}
if (warpOp.getResult(operandNumber).getType() == operand->get().getType()) {
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {extractOp.getSource()},
{extractOp.getSourceVectorType()}, newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value distributedVec = newWarpOp->getResult(newRetIndices[0]);
Value newExtract = vector::ExtractOp::create(
rewriter, loc, distributedVec, extractOp.getMixedPosition());
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber),
newExtract);
return success();
}
auto distributedType =
cast<VectorType>(warpOp.getResult(operandNumber).getType());
auto yieldedType = cast<VectorType>(operand->get().getType());
int64_t distributedDim = getDistributedDim(yieldedType, distributedType);
assert(distributedDim != -1 && "could not find distributed dimension");
(void)distributedDim;
SmallVector<int64_t> newDistributedShape(extractSrcType.getShape());
for (int i = 0; i < distributedType.getRank(); ++i)
newDistributedShape[i + extractOp.getNumIndices()] =
distributedType.getDimSize(i);
auto newDistributedType =
VectorType::get(newDistributedShape, distributedType.getElementType());
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {extractOp.getSource()}, {newDistributedType},
newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value distributedVec = newWarpOp->getResult(newRetIndices[0]);
Value newExtract = vector::ExtractOp::create(rewriter, loc, distributedVec,
extractOp.getMixedPosition());
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber),
newExtract);
return success();
}
};
struct WarpOpExtractScalar : public WarpDistributionPattern {
WarpOpExtractScalar(MLIRContext *ctx, WarpShuffleFromIdxFn fn,
PatternBenefit b = 1)
: WarpDistributionPattern(ctx, b), warpShuffleFromIdxFn(std::move(fn)) {}
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand =
getWarpResult(warpOp, llvm::IsaPred<vector::ExtractOp>);
if (!operand)
return failure();
unsigned int operandNumber = operand->getOperandNumber();
auto extractOp = operand->get().getDefiningOp<vector::ExtractOp>();
VectorType extractSrcType = extractOp.getSourceVectorType();
if (extractSrcType.getRank() > 1) {
return rewriter.notifyMatchFailure(
extractOp, "only 0-D or 1-D source supported for now");
}
if (!extractSrcType.getElementType().isF32() &&
!extractSrcType.getElementType().isInteger(32))
return rewriter.notifyMatchFailure(
extractOp, "only f32/i32 element types are supported");
bool is0dOrVec1Extract = extractSrcType.getNumElements() == 1;
Type elType = extractSrcType.getElementType();
VectorType distributedVecType;
if (!is0dOrVec1Extract) {
assert(extractSrcType.getRank() == 1 &&
"expected that extract src rank is 0 or 1");
if (extractSrcType.getShape()[0] % warpOp.getWarpSize() != 0)
return failure();
int64_t elementsPerLane =
extractSrcType.getShape()[0] / warpOp.getWarpSize();
distributedVecType = VectorType::get({elementsPerLane}, elType);
} else {
distributedVecType = extractSrcType;
}
SmallVector<Value> additionalResults{extractOp.getSource()};
SmallVector<Type> additionalResultTypes{distributedVecType};
additionalResults.append(
SmallVector<Value>(extractOp.getDynamicPosition()));
additionalResultTypes.append(
SmallVector<Type>(extractOp.getDynamicPosition().getTypes()));
Location loc = extractOp.getLoc();
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, additionalResults, additionalResultTypes,
newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value distributedVec = newWarpOp->getResult(newRetIndices[0]);
if (is0dOrVec1Extract) {
Value newExtract;
SmallVector<int64_t> indices(extractSrcType.getRank(), 0);
newExtract =
vector::ExtractOp::create(rewriter, loc, distributedVec, indices);
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber),
newExtract);
return success();
}
int64_t staticPos = extractOp.getStaticPosition()[0];
OpFoldResult pos = ShapedType::isDynamic(staticPos)
? (newWarpOp->getResult(newRetIndices[1]))
: OpFoldResult(rewriter.getIndexAttr(staticPos));
int64_t elementsPerLane = distributedVecType.getShape()[0];
AffineExpr sym0 = getAffineSymbolExpr(0, rewriter.getContext());
Value broadcastFromTid = affine::makeComposedAffineApply(
rewriter, loc, sym0.ceilDiv(elementsPerLane), pos);
Value newPos =
elementsPerLane == 1
? arith::ConstantIndexOp::create(rewriter, loc, 0).getResult()
: affine::makeComposedAffineApply(rewriter, loc,
sym0 % elementsPerLane, pos);
Value extracted =
vector::ExtractOp::create(rewriter, loc, distributedVec, newPos);
Value shuffled = warpShuffleFromIdxFn(
loc, rewriter, extracted, broadcastFromTid, newWarpOp.getWarpSize());
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber), shuffled);
return success();
}
private:
WarpShuffleFromIdxFn warpShuffleFromIdxFn;
};
struct WarpOpInsertScalar : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand = getWarpResult(warpOp, llvm::IsaPred<vector::InsertOp>);
if (!operand)
return failure();
unsigned int operandNumber = operand->getOperandNumber();
auto insertOp = operand->get().getDefiningOp<vector::InsertOp>();
VectorType vecType = insertOp.getDestVectorType();
VectorType distrType =
cast<VectorType>(warpOp.getResult(operandNumber).getType());
if (vecType.getRank() > 1) {
return rewriter.notifyMatchFailure(
insertOp, "only 0-D or 1-D source supported for now");
}
SmallVector<Value> additionalResults{insertOp.getDest(),
insertOp.getValueToStore()};
SmallVector<Type> additionalResultTypes{
distrType, insertOp.getValueToStore().getType()};
additionalResults.append(SmallVector<Value>(insertOp.getDynamicPosition()));
additionalResultTypes.append(
SmallVector<Type>(insertOp.getDynamicPosition().getTypes()));
Location loc = insertOp.getLoc();
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, additionalResults, additionalResultTypes,
newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value distributedVec = newWarpOp->getResult(newRetIndices[0]);
Value newSource = newWarpOp->getResult(newRetIndices[1]);
rewriter.setInsertionPointAfter(newWarpOp);
OpFoldResult pos;
if (vecType.getRank() != 0) {
int64_t staticPos = insertOp.getStaticPosition()[0];
pos = ShapedType::isDynamic(staticPos)
? (newWarpOp->getResult(newRetIndices[2]))
: OpFoldResult(rewriter.getIndexAttr(staticPos));
}
if (vecType == distrType) {
Value newInsert;
SmallVector<OpFoldResult> indices;
if (pos) {
indices.push_back(pos);
}
newInsert = vector::InsertOp::create(rewriter, loc, newSource,
distributedVec, indices);
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber),
newInsert);
return success();
}
int64_t elementsPerLane = distrType.getShape()[0];
AffineExpr sym0 = getAffineSymbolExpr(0, rewriter.getContext());
Value insertingLane = affine::makeComposedAffineApply(
rewriter, loc, sym0.ceilDiv(elementsPerLane), pos);
OpFoldResult newPos = affine::makeComposedFoldedAffineApply(
rewriter, loc, sym0 % elementsPerLane, pos);
Value isInsertingLane =
arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq,
newWarpOp.getLaneid(), insertingLane);
Value newResult =
scf::IfOp::create(
rewriter, loc, isInsertingLane,
[&](OpBuilder &builder, Location loc) {
Value newInsert = vector::InsertOp::create(
builder, loc, newSource, distributedVec, newPos);
scf::YieldOp::create(builder, loc, newInsert);
},
[&](OpBuilder &builder, Location loc) {
scf::YieldOp::create(builder, loc, distributedVec);
})
.getResult(0);
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber), newResult);
return success();
}
};
struct WarpOpInsert : public WarpDistributionPattern {
using Base::Base;
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *operand = getWarpResult(warpOp, llvm::IsaPred<vector::InsertOp>);
if (!operand)
return failure();
unsigned int operandNumber = operand->getOperandNumber();
auto insertOp = operand->get().getDefiningOp<vector::InsertOp>();
Location loc = insertOp.getLoc();
if (insertOp.getDestVectorType().getRank() <= 1) {
return failure();
}
if (warpOp.getResult(operandNumber).getType() == operand->get().getType()) {
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {insertOp.getValueToStore(), insertOp.getDest()},
{insertOp.getValueToStoreType(), insertOp.getDestVectorType()},
newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value distributedSrc = newWarpOp->getResult(newRetIndices[0]);
Value distributedDest = newWarpOp->getResult(newRetIndices[1]);
Value newResult = vector::InsertOp::create(rewriter, loc, distributedSrc,
distributedDest,
insertOp.getMixedPosition());
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber),
newResult);
return success();
}
auto distrDestType =
cast<VectorType>(warpOp.getResult(operandNumber).getType());
auto yieldedType = cast<VectorType>(operand->get().getType());
int64_t distrDestDim = -1;
for (int64_t i = 0; i < yieldedType.getRank(); ++i) {
if (distrDestType.getDimSize(i) != yieldedType.getDimSize(i)) {
assert(distrDestDim == -1 && "found multiple distributed dims");
distrDestDim = i;
}
}
assert(distrDestDim != -1 && "could not find distributed dimension");
VectorType srcVecType = cast<VectorType>(insertOp.getValueToStoreType());
SmallVector<int64_t> distrSrcShape(srcVecType.getShape());
int64_t distrSrcDim = distrDestDim - insertOp.getNumIndices();
if (distrSrcDim >= 0)
distrSrcShape[distrSrcDim] = distrDestType.getDimSize(distrDestDim);
auto distrSrcType =
VectorType::get(distrSrcShape, distrDestType.getElementType());
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, {insertOp.getValueToStore(), insertOp.getDest()},
{distrSrcType, distrDestType}, newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value distributedSrc = newWarpOp->getResult(newRetIndices[0]);
Value distributedDest = newWarpOp->getResult(newRetIndices[1]);
Value newResult;
if (distrSrcDim >= 0) {
newResult = vector::InsertOp::create(rewriter, loc, distributedSrc,
distributedDest,
insertOp.getMixedPosition());
} else {
int64_t elementsPerLane = distrDestType.getDimSize(distrDestDim);
SmallVector<OpFoldResult> pos = insertOp.getMixedPosition();
SmallVector<int64_t> newPos = getAsIntegers(pos);
Value insertingLane = arith::ConstantIndexOp::create(
rewriter, loc, newPos[distrDestDim] / elementsPerLane);
Value isInsertingLane =
arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq,
newWarpOp.getLaneid(), insertingLane);
newPos[distrDestDim] %= elementsPerLane;
auto insertingBuilder = [&](OpBuilder &builder, Location loc) {
Value newInsert = vector::InsertOp::create(builder, loc, distributedSrc,
distributedDest, newPos);
scf::YieldOp::create(builder, loc, newInsert);
};
auto nonInsertingBuilder = [&](OpBuilder &builder, Location loc) {
scf::YieldOp::create(builder, loc, distributedDest);
};
newResult = scf::IfOp::create(rewriter, loc, isInsertingLane,
insertingBuilder,
nonInsertingBuilder)
.getResult(0);
}
rewriter.replaceAllUsesWith(newWarpOp->getResult(operandNumber), newResult);
return success();
}
};
struct WarpOpScfIfOp : public WarpDistributionPattern {
WarpOpScfIfOp(MLIRContext *ctx, DistributionMapFn fn, PatternBenefit b = 1)
: WarpDistributionPattern(ctx, b), distributionMapFn(std::move(fn)) {}
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
gpu::YieldOp warpOpYield = warpOp.getTerminator();
Operation *lastNode = warpOpYield->getPrevNode();
auto ifOp = dyn_cast_or_null<scf::IfOp>(lastNode);
if (!ifOp)
return failure();
SmallVector<Value> nonIfYieldValues;
SmallVector<unsigned> nonIfYieldIndices;
llvm::SmallDenseMap<unsigned, unsigned> ifResultMapping;
llvm::SmallDenseMap<unsigned, VectorType> ifResultDistTypes;
for (OpOperand &yieldOperand : warpOpYield->getOpOperands()) {
const unsigned yieldOperandIdx = yieldOperand.getOperandNumber();
if (yieldOperand.get().getDefiningOp() != ifOp.getOperation()) {
nonIfYieldValues.push_back(yieldOperand.get());
nonIfYieldIndices.push_back(yieldOperandIdx);
continue;
}
OpResult ifResult = cast<OpResult>(yieldOperand.get());
const unsigned ifResultIdx = ifResult.getResultNumber();
ifResultMapping[yieldOperandIdx] = ifResultIdx;
if (!isa<VectorType>(ifResult.getType()))
continue;
VectorType distType =
cast<VectorType>(warpOp.getResult(yieldOperandIdx).getType());
ifResultDistTypes[ifResultIdx] = distType;
}
auto [escapingValuesThen, escapingValueInputTypesThen,
escapingValueDistTypesThen] =
getInnerRegionEscapingValues(warpOp, ifOp.getThenRegion(),
distributionMapFn);
auto [escapingValuesElse, escapingValueInputTypesElse,
escapingValueDistTypesElse] =
getInnerRegionEscapingValues(warpOp, ifOp.getElseRegion(),
distributionMapFn);
if (llvm::is_contained(escapingValueDistTypesThen, Type{}) ||
llvm::is_contained(escapingValueDistTypesElse, Type{}))
return failure();
SmallVector<Value> newWarpOpYieldValues{ifOp.getCondition()};
newWarpOpYieldValues.append(escapingValuesThen.begin(),
escapingValuesThen.end());
newWarpOpYieldValues.append(escapingValuesElse.begin(),
escapingValuesElse.end());
SmallVector<Type> newWarpOpDistTypes{ifOp.getCondition().getType()};
newWarpOpDistTypes.append(escapingValueDistTypesThen.begin(),
escapingValueDistTypesThen.end());
newWarpOpDistTypes.append(escapingValueDistTypesElse.begin(),
escapingValueDistTypesElse.end());
for (auto [idx, val] :
llvm::zip_equal(nonIfYieldIndices, nonIfYieldValues)) {
newWarpOpYieldValues.push_back(val);
newWarpOpDistTypes.push_back(warpOp.getResult(idx).getType());
}
SmallVector<size_t> newIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, newWarpOpYieldValues, newWarpOpDistTypes, newIndices);
SmallVector<Type> newIfOpDistResTypes;
for (auto [i, res] : llvm::enumerate(ifOp.getResults())) {
Type distType = cast<Value>(res).getType();
if (auto vecType = dyn_cast<VectorType>(distType)) {
AffineMap map = distributionMapFn(cast<Value>(res));
distType = ifResultDistTypes.count(i)
? ifResultDistTypes[i]
: getDistributedType(vecType, map, warpOp.getWarpSize());
}
newIfOpDistResTypes.push_back(distType);
}
OpBuilder::InsertionGuard g(rewriter);
rewriter.setInsertionPointAfter(newWarpOp);
auto newIfOp = scf::IfOp::create(
rewriter, ifOp.getLoc(), newIfOpDistResTypes,
newWarpOp.getResult(newIndices[0]), static_cast<bool>(ifOp.thenBlock()),
static_cast<bool>(ifOp.elseBlock()));
auto encloseRegionInWarpOp =
[&](Block *oldIfBranch, Block *newIfBranch,
llvm::SmallSetVector<Value, 32> &escapingValues,
SmallVector<Type> &escapingValueInputTypes,
size_t warpResRangeStart) {
OpBuilder::InsertionGuard g(rewriter);
if (!newIfBranch)
return;
rewriter.setInsertionPointToStart(newIfBranch);
llvm::SmallDenseMap<Value, int64_t> escapeValToBlockArgIndex;
SmallVector<Value> innerWarpInputVals;
SmallVector<Type> innerWarpInputTypes;
for (size_t i = 0; i < escapingValues.size();
++i, ++warpResRangeStart) {
innerWarpInputVals.push_back(
newWarpOp.getResult(newIndices[warpResRangeStart]));
escapeValToBlockArgIndex[escapingValues[i]] =
innerWarpInputTypes.size();
innerWarpInputTypes.push_back(escapingValueInputTypes[i]);
}
auto innerWarp = WarpExecuteOnLane0Op::create(
rewriter, newWarpOp.getLoc(), newIfOp.getResultTypes(),
newWarpOp.getLaneid(), newWarpOp.getWarpSize(),
innerWarpInputVals, innerWarpInputTypes);
innerWarp.getWarpRegion().takeBody(*oldIfBranch->getParent());
innerWarp.getWarpRegion().addArguments(
innerWarpInputTypes,
SmallVector<Location>(innerWarpInputTypes.size(), ifOp.getLoc()));
SmallVector<Value> yieldOperands;
for (Value operand : oldIfBranch->getTerminator()->getOperands())
yieldOperands.push_back(operand);
rewriter.eraseOp(oldIfBranch->getTerminator());
rewriter.setInsertionPointToEnd(innerWarp.getBody());
gpu::YieldOp::create(rewriter, innerWarp.getLoc(), yieldOperands);
rewriter.setInsertionPointAfter(innerWarp);
scf::YieldOp::create(rewriter, ifOp.getLoc(), innerWarp.getResults());
innerWarp.walk([&](Operation *op) {
for (OpOperand &operand : op->getOpOperands()) {
auto it = escapeValToBlockArgIndex.find(operand.get());
if (it == escapeValToBlockArgIndex.end())
continue;
operand.set(innerWarp.getBodyRegion().getArgument(it->second));
}
});
mlir::vector::moveScalarUniformCode(innerWarp);
};
encloseRegionInWarpOp(&ifOp.getThenRegion().front(),
&newIfOp.getThenRegion().front(), escapingValuesThen,
escapingValueInputTypesThen, 1);
if (!ifOp.getElseRegion().empty())
encloseRegionInWarpOp(&ifOp.getElseRegion().front(),
&newIfOp.getElseRegion().front(),
escapingValuesElse, escapingValueInputTypesElse,
1 + escapingValuesThen.size());
for (auto [origIdx, newIdx] : ifResultMapping)
rewriter.replaceAllUsesExcept(newWarpOp.getResult(origIdx),
newIfOp.getResult(newIdx), newIfOp);
return success();
}
private:
DistributionMapFn distributionMapFn;
};
struct WarpOpScfForOp : public WarpDistributionPattern {
WarpOpScfForOp(MLIRContext *ctx, DistributionMapFn fn, PatternBenefit b = 1)
: WarpDistributionPattern(ctx, b), distributionMapFn(std::move(fn)) {}
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
gpu::YieldOp warpOpYield = warpOp.getTerminator();
Operation *lastNode = warpOpYield->getPrevNode();
auto forOp = dyn_cast_or_null<scf::ForOp>(lastNode);
if (!forOp)
return failure();
auto [escapingValues, escapingValueInputTypes, escapingValueDistTypes] =
getInnerRegionEscapingValues(warpOp, forOp.getBodyRegion(),
distributionMapFn);
if (llvm::is_contained(escapingValueDistTypes, Type{}))
return failure();
SmallVector<Value> nonForYieldedValues;
SmallVector<unsigned> nonForResultIndices;
llvm::SmallDenseMap<unsigned, unsigned> forResultMapping;
llvm::SmallDenseMap<unsigned, VectorType> forResultDistTypes;
for (OpOperand &yieldOperand : warpOpYield->getOpOperands()) {
if (yieldOperand.get().getDefiningOp() != forOp.getOperation()) {
nonForYieldedValues.push_back(yieldOperand.get());
nonForResultIndices.push_back(yieldOperand.getOperandNumber());
continue;
}
OpResult forResult = cast<OpResult>(yieldOperand.get());
unsigned int forResultNumber = forResult.getResultNumber();
forResultMapping[yieldOperand.getOperandNumber()] = forResultNumber;
if (!isa<VectorType>(forResult.getType()))
continue;
VectorType distType = cast<VectorType>(
warpOp.getResult(yieldOperand.getOperandNumber()).getType());
forResultDistTypes[forResultNumber] = distType;
}
SmallVector<Value> newWarpOpYieldValues;
SmallVector<Type> newWarpOpDistTypes;
newWarpOpYieldValues.insert(
newWarpOpYieldValues.end(),
{forOp.getLowerBound(), forOp.getUpperBound(), forOp.getStep()});
newWarpOpDistTypes.insert(newWarpOpDistTypes.end(),
{forOp.getLowerBound().getType(),
forOp.getUpperBound().getType(),
forOp.getStep().getType()});
for (auto [i, initArg] : llvm::enumerate(forOp.getInitArgs())) {
newWarpOpYieldValues.push_back(initArg);
Type distType = initArg.getType();
if (auto vecType = dyn_cast<VectorType>(distType)) {
AffineMap map = distributionMapFn(initArg);
distType = forResultDistTypes.count(i)
? forResultDistTypes[i]
: getDistributedType(vecType, map, warpOp.getWarpSize());
}
newWarpOpDistTypes.push_back(distType);
}
newWarpOpYieldValues.insert(newWarpOpYieldValues.end(),
escapingValues.begin(), escapingValues.end());
newWarpOpDistTypes.insert(newWarpOpDistTypes.end(),
escapingValueDistTypes.begin(),
escapingValueDistTypes.end());
for (auto [i, v] :
llvm::zip_equal(nonForResultIndices, nonForYieldedValues)) {
newWarpOpYieldValues.push_back(v);
newWarpOpDistTypes.push_back(warpOp.getResult(i).getType());
}
SmallVector<size_t> newIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, newWarpOpYieldValues, newWarpOpDistTypes, newIndices);
const unsigned initArgsStartIdx = 3;
const unsigned escapingValuesStartIdx =
initArgsStartIdx +
forOp.getInitArgs().size();
SmallVector<Value> newForOpOperands;
for (size_t i = initArgsStartIdx; i < escapingValuesStartIdx; ++i)
newForOpOperands.push_back(newWarpOp.getResult(newIndices[i]));
OpBuilder::InsertionGuard g(rewriter);
rewriter.setInsertionPointAfter(newWarpOp);
auto newForOp = scf::ForOp::create(
rewriter, forOp.getLoc(),
newWarpOp.getResult(newIndices[0]),
newWarpOp.getResult(newIndices[1]),
newWarpOp.getResult(newIndices[2]), newForOpOperands,
nullptr, forOp.getUnsignedCmp());
rewriter.setInsertionPointToStart(newForOp.getBody());
SmallVector<Value> innerWarpInput(newForOp.getRegionIterArgs().begin(),
newForOp.getRegionIterArgs().end());
SmallVector<Type> innerWarpInputType(forOp.getResultTypes().begin(),
forOp.getResultTypes().end());
llvm::SmallDenseMap<Value, int64_t> argIndexMapping;
for (size_t i = escapingValuesStartIdx;
i < escapingValuesStartIdx + escapingValues.size(); ++i) {
innerWarpInput.push_back(newWarpOp.getResult(newIndices[i]));
argIndexMapping[escapingValues[i - escapingValuesStartIdx]] =
innerWarpInputType.size();
innerWarpInputType.push_back(
escapingValueInputTypes[i - escapingValuesStartIdx]);
}
auto innerWarp = WarpExecuteOnLane0Op::create(
rewriter, newWarpOp.getLoc(), newForOp.getResultTypes(),
newWarpOp.getLaneid(), newWarpOp.getWarpSize(), innerWarpInput,
innerWarpInputType);
SmallVector<Value> argMapping;
argMapping.push_back(newForOp.getInductionVar());
for (Value args : innerWarp.getBody()->getArguments())
argMapping.push_back(args);
argMapping.resize(forOp.getBody()->getNumArguments());
SmallVector<Value> yieldOperands;
for (Value operand : forOp.getBody()->getTerminator()->getOperands())
yieldOperands.push_back(operand);
rewriter.eraseOp(forOp.getBody()->getTerminator());
rewriter.mergeBlocks(forOp.getBody(), innerWarp.getBody(), argMapping);
rewriter.setInsertionPointToEnd(innerWarp.getBody());
gpu::YieldOp::create(rewriter, innerWarp.getLoc(), yieldOperands);
rewriter.setInsertionPointAfter(innerWarp);
if (!innerWarp.getResults().empty())
scf::YieldOp::create(rewriter, forOp.getLoc(), innerWarp.getResults());
for (auto [origIdx, newIdx] : forResultMapping)
rewriter.replaceAllUsesExcept(newWarpOp.getResult(origIdx),
newForOp.getResult(newIdx), newForOp);
newForOp.walk([&](Operation *op) {
for (OpOperand &operand : op->getOpOperands()) {
auto it = argIndexMapping.find(operand.get());
if (it == argIndexMapping.end())
continue;
operand.set(innerWarp.getBodyRegion().getArgument(it->second));
}
});
mlir::vector::moveScalarUniformCode(innerWarp);
return success();
}
private:
DistributionMapFn distributionMapFn;
};
struct WarpOpReduction : public WarpDistributionPattern {
WarpOpReduction(MLIRContext *context,
DistributedReductionFn distributedReductionFn,
PatternBenefit benefit = 1)
: WarpDistributionPattern(context, benefit),
distributedReductionFn(std::move(distributedReductionFn)) {}
LogicalResult matchAndRewrite(WarpExecuteOnLane0Op warpOp,
PatternRewriter &rewriter) const override {
OpOperand *yieldOperand =
getWarpResult(warpOp, llvm::IsaPred<vector::ReductionOp>);
if (!yieldOperand)
return failure();
auto reductionOp =
cast<vector::ReductionOp>(yieldOperand->get().getDefiningOp());
auto vectorType = cast<VectorType>(reductionOp.getVector().getType());
if (vectorType.getRank() != 1)
return rewriter.notifyMatchFailure(
warpOp, "Only rank 1 reductions can be distributed.");
if (vectorType.getShape()[0] % warpOp.getWarpSize() != 0)
return rewriter.notifyMatchFailure(
warpOp, "Reduction vector dimension must match was size.");
if (!reductionOp.getType().isIntOrFloat())
return rewriter.notifyMatchFailure(
warpOp, "Reduction distribution currently only supports floats and "
"integer types.");
int64_t numElements = vectorType.getShape()[0] / warpOp.getWarpSize();
unsigned operandIndex = yieldOperand->getOperandNumber();
SmallVector<Value> yieldValues = {reductionOp.getVector()};
SmallVector<Type> retTypes = {
VectorType::get({numElements}, reductionOp.getType())};
if (reductionOp.getAcc()) {
yieldValues.push_back(reductionOp.getAcc());
retTypes.push_back(reductionOp.getAcc().getType());
}
SmallVector<size_t> newRetIndices;
WarpExecuteOnLane0Op newWarpOp = moveRegionToNewWarpOpAndAppendReturns(
rewriter, warpOp, yieldValues, retTypes, newRetIndices);
rewriter.setInsertionPointAfter(newWarpOp);
Value laneValVec = newWarpOp.getResult(newRetIndices[0]);
Value fullReduce =
distributedReductionFn(reductionOp.getLoc(), rewriter, laneValVec,
reductionOp.getKind(), newWarpOp.getWarpSize());
if (reductionOp.getAcc()) {
fullReduce = vector::makeArithReduction(
rewriter, reductionOp.getLoc(), reductionOp.getKind(), fullReduce,
newWarpOp.getResult(newRetIndices[1]));
}
rewriter.replaceAllUsesWith(newWarpOp.getResult(operandIndex), fullReduce);
return success();
}
private:
DistributedReductionFn distributedReductionFn;
};
}
void mlir::vector::populateWarpExecuteOnLane0OpToScfForPattern(
RewritePatternSet &patterns,
const WarpExecuteOnLane0LoweringOptions &options, PatternBenefit benefit) {
patterns.add<WarpOpToScfIfPattern>(patterns.getContext(), options, benefit);
}
void mlir::vector::populateDistributeTransferWriteOpPatterns(
RewritePatternSet &patterns, const DistributionMapFn &distributionMapFn,
unsigned maxNumElementsToExtract, PatternBenefit benefit) {
patterns.add<WarpOpTransferWrite>(patterns.getContext(), distributionMapFn,
maxNumElementsToExtract, benefit);
}
void mlir::vector::populatePropagateWarpVectorDistributionPatterns(
RewritePatternSet &patterns, const DistributionMapFn &distributionMapFn,
const WarpShuffleFromIdxFn &warpShuffleFromIdxFn, PatternBenefit benefit,
PatternBenefit readBenefit) {
patterns.add<WarpOpTransferRead>(patterns.getContext(), readBenefit);
patterns
.add<WarpOpElementwise, WarpOpDeadResult, WarpOpBroadcast,
WarpOpShapeCast, WarpOpExtract, WarpOpForwardOperand, WarpOpConstant,
WarpOpInsertScalar, WarpOpInsert, WarpOpCreateMask,
WarpOpExtractStridedSlice, WarpOpInsertStridedSlice, WarpOpStep>(
patterns.getContext(), benefit);
patterns.add<WarpOpExtractScalar>(patterns.getContext(), warpShuffleFromIdxFn,
benefit);
patterns.add<WarpOpScfForOp>(patterns.getContext(), distributionMapFn,
benefit);
patterns.add<WarpOpScfIfOp>(patterns.getContext(), distributionMapFn,
benefit);
}
void mlir::vector::populateDistributeReduction(
RewritePatternSet &patterns,
const DistributedReductionFn &distributedReductionFn,
PatternBenefit benefit) {
patterns.add<WarpOpReduction>(patterns.getContext(), distributedReductionFn,
benefit);
}
static bool canBeHoisted(Operation *op,
function_ref<bool(Value)> definedOutside) {
return llvm::all_of(op->getOperands(), definedOutside) &&
isMemoryEffectFree(op) && op->getNumRegions() == 0;
}
void mlir::vector::moveScalarUniformCode(WarpExecuteOnLane0Op warpOp) {
Block *body = warpOp.getBody();
llvm::SmallSetVector<Operation *, 8> opsToMove;
auto isDefinedOutsideOfBody = [&](Value value) {
auto *definingOp = value.getDefiningOp();
return (definingOp && opsToMove.count(definingOp)) ||
warpOp.isDefinedOutsideOfRegion(value);
};
for (auto &op : body->without_terminator()) {
bool hasVectorResult = llvm::any_of(op.getResults(), [](Value result) {
return isa<VectorType>(result.getType());
});
if (!hasVectorResult && canBeHoisted(&op, isDefinedOutsideOfBody))
opsToMove.insert(&op);
}
for (Operation *op : opsToMove)
op->moveBefore(warpOp);
}