36#include "llvm/ADT/STLExtras.h"
37#include "llvm/ADT/TypeSwitch.h"
38#include "llvm/Support/CommandLine.h"
39#include "llvm/Support/ErrorHandling.h"
40#include "llvm/Support/FormatVariadic.h"
41#include "llvm/Support/InterleavedRange.h"
42#include "llvm/Support/StringSaver.h"
50#include "mlir/Dialect/GPU/IR/GPUOpsDialect.cpp.inc"
56int64_t GPUBlockMappingAttr::getMappingId()
const {
57 return static_cast<int64_t>(getBlock());
60bool GPUBlockMappingAttr::isLinearMapping()
const {
61 return getMappingId() >=
static_cast<int64_t>(MappingId::LinearDim0);
64int64_t GPUBlockMappingAttr::getRelativeIndex()
const {
65 return isLinearMapping()
66 ? getMappingId() -
static_cast<int64_t>(MappingId::LinearDim0)
70int64_t GPUWarpgroupMappingAttr::getMappingId()
const {
71 return static_cast<int64_t>(getWarpgroup());
74bool GPUWarpgroupMappingAttr::isLinearMapping()
const {
75 return getMappingId() >=
static_cast<int64_t>(MappingId::LinearDim0);
78int64_t GPUWarpgroupMappingAttr::getRelativeIndex()
const {
79 return isLinearMapping()
80 ? getMappingId() -
static_cast<int64_t>(MappingId::LinearDim0)
84int64_t GPUWarpMappingAttr::getMappingId()
const {
85 return static_cast<int64_t>(getWarp());
88bool GPUWarpMappingAttr::isLinearMapping()
const {
89 return getMappingId() >=
static_cast<int64_t>(MappingId::LinearDim0);
92int64_t GPUWarpMappingAttr::getRelativeIndex()
const {
93 return isLinearMapping()
94 ? getMappingId() -
static_cast<int64_t>(MappingId::LinearDim0)
98int64_t GPUThreadMappingAttr::getMappingId()
const {
99 return static_cast<int64_t>(getThread());
102bool GPUThreadMappingAttr::isLinearMapping()
const {
103 return getMappingId() >=
static_cast<int64_t>(MappingId::LinearDim0);
106int64_t GPUThreadMappingAttr::getRelativeIndex()
const {
107 return isLinearMapping()
108 ? getMappingId() -
static_cast<int64_t>(MappingId::LinearDim0)
112int64_t GPULaneMappingAttr::getMappingId()
const {
113 return static_cast<int64_t>(getLane());
116bool GPULaneMappingAttr::isLinearMapping()
const {
117 return getMappingId() >=
static_cast<int64_t>(MappingId::LinearDim0);
120int64_t GPULaneMappingAttr::getRelativeIndex()
const {
121 return isLinearMapping()
122 ? getMappingId() -
static_cast<int64_t>(MappingId::LinearDim0)
126int64_t GPUMappingMaskAttr::getMaxNumPhysicalIds()
const {
return 64; }
138Value GPUMappingMaskAttr::createLogicalLinearMappingId(
142 arith::ConstantOp::create(
b, loc,
b.getI64IntegerAttr(getMask()));
143 Value one = arith::ConstantOp::create(
b, loc,
b.getI64IntegerAttr(1));
144 Value filter = arith::ShLIOp::create(
b, loc, one, physicalLinearMappingId);
145 filter = arith::SubIOp::create(
b, loc, filter, one);
146 Value filteredId = arith::AndIOp::create(
b, loc, mask, filter);
147 return math::CtPopOp::create(
b, loc, filteredId);
160Value GPUMappingMaskAttr::createIsActiveIdPredicate(
164 arith::ConstantOp::create(
b, loc,
b.getI64IntegerAttr(getMask()));
165 Value one = arith::ConstantOp::create(
b, loc,
b.getI64IntegerAttr(1));
166 Value filter = arith::ShLIOp::create(
b, loc, one, physicalLinearMappingId);
167 Value filtered = arith::AndIOp::create(
b, loc, mask, filter);
168 Value zero = arith::ConstantOp::create(
b, loc,
b.getI64IntegerAttr(0));
169 return arith::CmpIOp::create(
b, loc, arith::CmpIPredicate::ne, filtered,
173int64_t GPUMemorySpaceMappingAttr::getMappingId()
const {
174 return static_cast<int64_t>(getAddressSpace());
177bool GPUMemorySpaceMappingAttr::isLinearMapping()
const {
178 llvm_unreachable(
"GPUMemorySpaceMappingAttr does not support linear mapping");
181int64_t GPUMemorySpaceMappingAttr::getRelativeIndex()
const {
182 llvm_unreachable(
"GPUMemorySpaceMappingAttr does not support relative index");
199 elementType, operand);
213 return elementType.
isF16() || elementType.
isF32() || elementType.
isF64() ||
222 if (operand !=
"AOp" && operand !=
"BOp" && operand !=
"COp")
223 return emitError() <<
"operand expected to be one of AOp, BOp or COp";
225 if (
shape.size() != 2)
226 return emitError() <<
"MMAMatrixType must have exactly two dimensions";
230 <<
"MMAMatrixType elements must be SI8, UI8, I32, F16, F32, or F64";
239bool GPUDialect::isWorkgroupMemoryAddressSpace(
Attribute memorySpace) {
242 if (
auto gpuAttr = llvm::dyn_cast<gpu::AddressSpaceAttr>(memorySpace))
243 return gpuAttr.getValue() == getWorkgroupAddressSpace();
247bool GPUDialect::hasWorkgroupMemoryAddressSpace(MemRefType type) {
248 Attribute memorySpace = type.getMemorySpace();
249 return isWorkgroupMemoryAddressSpace(memorySpace);
252bool GPUDialect::isConstantMemoryAddressSpace(
Attribute memorySpace) {
255 if (
auto gpuAttr = llvm::dyn_cast<gpu::AddressSpaceAttr>(memorySpace))
256 return gpuAttr.getValue() == getConstantAddressSpace();
260bool GPUDialect::hasConstantMemoryAddressSpace(MemRefType type) {
261 Attribute memorySpace = type.getMemorySpace();
262 return isConstantMemoryAddressSpace(memorySpace);
265bool GPUDialect::isKernel(Operation *op) {
266 if (
auto gpuFunc = dyn_cast<GPUFuncOp>(op))
267 return gpuFunc.isKernel();
268 return static_cast<bool>(
275struct GPUInlinerInterface :
public DialectInlinerInterface {
276 using DialectInlinerInterface::DialectInlinerInterface;
279 bool isLegalToInline(Operation *, Region *,
bool, IRMapping &)
const final {
285void GPUDialect::initialize() {
286 addTypes<AsyncTokenType>();
287 addTypes<MMAMatrixType>();
288 addTypes<NamedBarrierType>();
289 addTypes<SparseDnTensorHandleType>();
290 addTypes<SparseSpMatHandleType>();
291 addTypes<SparseSpGEMMOpHandleType>();
294#include "mlir/Dialect/GPU/IR/GPUOps.cpp.inc"
297#define GET_ATTRDEF_LIST
298#include "mlir/Dialect/GPU/IR/GPUOpsAttributes.cpp.inc"
300 addInterfaces<GPUInlinerInterface>();
301 declarePromisedInterface<bufferization::BufferDeallocationOpInterface,
303 declarePromisedInterfaces<ValueBoundsOpInterface, ClusterDimOp,
304 ClusterDimBlocksOp, ClusterIdOp, ClusterBlockIdOp,
305 BlockDimOp, BlockIdOp, GridDimOp, ThreadIdOp,
306 LaneIdOp, SubgroupIdOp, GlobalIdOp, NumSubgroupsOp,
307 SubgroupSizeOp, LaunchOp, SubgroupBroadcastOp>();
308 declarePromisedInterfaces<memref::IndexedAccessOpInterface,
309 SubgroupMmaLoadMatrixOp,
310 SubgroupMmaStoreMatrixOp>();
316 return "sparse.dntensor_handle";
318 return "sparse.spmat_handle";
320 return "sparse.spgemmop_handle";
322 llvm_unreachable(
"unknown sparse handle kind");
326Type GPUDialect::parseType(DialectAsmParser &parser)
const {
334 if (keyword ==
"async.token")
337 if (keyword ==
"mma_matrix") {
345 SmallVector<int64_t> shape;
366 shape, elementType, operand);
369 if (keyword ==
"named_barrier")
384void GPUDialect::printType(Type type, DialectAsmPrinter &os)
const {
387 .Case<NamedBarrierType>([&](Type) { os <<
"named_barrier"; })
388 .Case<SparseDnTensorHandleType>([&](Type) {
391 .Case<SparseSpMatHandleType>(
393 .Case<SparseSpGEMMOpHandleType>([&](Type) {
399 for (
auto dim = shape.begin(), e = shape.end() - 1; dim != e; ++dim)
402 os <<
", \"" << fragTy.
getOperand() <<
"\"" <<
'>';
404 .DefaultUnreachable(
"unexpected 'gpu' type kind");
409 auto array = dyn_cast<DenseI32ArrayAttr>(attr.
getValue());
412 " must be a dense i32 array");
413 if (array.size() != 3)
415 " must contain exactly 3 elements");
419LogicalResult GPUDialect::verifyOperationAttribute(Operation *op,
420 NamedAttribute attr) {
421 if (attr.
getName() == getKnownBlockSizeAttrHelper().getName())
423 if (attr.
getName() == getKnownGridSizeAttrHelper().getName())
425 if (attr.
getName() == getKnownClusterSizeAttrHelper().getName())
427 if (!llvm::isa<UnitAttr>(attr.
getValue()) ||
428 attr.
getName() != getContainerModuleAttrName())
431 auto module = dyn_cast<ModuleOp>(op);
434 << getContainerModuleAttrName() <<
"' attribute to be attached to '"
435 << ModuleOp::getOperationName() <<
'\'';
449 return parser.
emitError(loc,
"needs to be named when marked 'async'");
464 if (asyncDependencies.empty())
468 printer << llvm::interleaved_array(asyncDependencies);
496 p <<
' ' << keyword <<
'(';
497 llvm::interleaveComma(
498 llvm::enumerate(values), p, [&p, attributes](
auto pair) {
499 BlockArgument v = pair.value();
500 p << v <<
" : " << v.
getType();
502 size_t attributionIndex = pair.index();
503 DictionaryAttr attrs;
504 if (attributes && attributionIndex < attributes.size())
505 attrs = llvm::cast<DictionaryAttr>(attributes[attributionIndex]);
515 gpu::AddressSpace memorySpace) {
516 for (
Value v : attributions) {
517 auto type = llvm::dyn_cast<MemRefType>(v.
getType());
519 return op->
emitOpError() <<
"expected memref type in attribution";
524 llvm::dyn_cast_or_null<gpu::AddressSpaceAttr>(type.getMemorySpace());
527 if (addressSpace.getValue() != memorySpace)
529 <<
"expected memory space " << stringifyAddressSpace(memorySpace)
530 <<
" in attribution";
541 using Kind = gpu::AllReduceOperation;
542 if (llvm::is_contained(
543 {Kind::MINNUMF, Kind::MAXNUMF, Kind::MINIMUMF, Kind::MAXIMUMF},
545 if (!isa<FloatType>(resType))
549 if (llvm::is_contained({Kind::MINSI, Kind::MINUI, Kind::MAXSI, Kind::MAXUI,
550 Kind::AND, Kind::OR, Kind::XOR},
552 if (!isa<IntegerType>(resType))
559LogicalResult gpu::AllReduceOp::verifyRegions() {
560 if (getBody().empty() != getOp().has_value())
561 return emitError(
"expected either an op attribute or a non-empty body");
562 if (!getBody().empty()) {
563 if (getBody().getNumArguments() != 2)
564 return emitError(
"expected two region arguments");
565 for (
auto argument : getBody().getArguments()) {
566 if (argument.getType() !=
getType())
567 return emitError(
"incorrect region argument type");
569 unsigned yieldCount = 0;
570 for (
Block &block : getBody()) {
571 if (
auto yield = dyn_cast<gpu::YieldOp>(block.getTerminator())) {
572 if (yield.getNumOperands() != 1)
573 return emitError(
"expected one gpu.yield operand");
574 if (yield.getOperand(0).getType() !=
getType())
575 return emitError(
"incorrect gpu.yield type");
580 return emitError(
"expected gpu.yield op in region");
582 gpu::AllReduceOperation opName = *getOp();
584 return emitError() <<
'`' << gpu::stringifyAllReduceOperation(opName)
585 <<
"` reduction operation is not compatible with type "
594 auto launchOp = dyn_cast<gpu::LaunchOp>(op->
getParentOp());
598 Region &body = launchOp.getBody();
599 assert(!body.
empty() &&
"Invalid region");
605OpFoldResult gpu::AllReduceOp::fold(FoldAdaptor ) {
618LogicalResult gpu::SubgroupReduceOp::verify() {
620 if (
auto vecTy = dyn_cast<VectorType>(elemType)) {
621 if (vecTy.isScalable())
622 return emitOpError() <<
"is not compatible with scalable vector types";
624 elemType = vecTy.getElementType();
627 gpu::AllReduceOperation opName = getOp();
629 return emitError() <<
'`' << gpu::stringifyAllReduceOperation(opName)
630 <<
"` reduction operation is not compatible with type "
634 auto clusterSize = getClusterSize();
636 uint32_t size = *clusterSize;
637 if (!llvm::isPowerOf2_32(size)) {
638 return emitOpError() <<
"cluster size " << size
639 <<
" is not a power of two";
643 uint32_t stride = getClusterStride();
644 if (stride != 1 && !clusterSize) {
645 return emitOpError() <<
"cluster stride can only be specified if cluster "
648 if (!llvm::isPowerOf2_32(stride)) {
649 return emitOpError() <<
"cluster stride " << stride
650 <<
" is not a power of two";
656OpFoldResult gpu::SubgroupReduceOp::fold(FoldAdaptor ) {
657 if (getClusterSize() == 1)
674 if (!op->template hasTrait<OpTrait::AttrSizedOperandSegments>())
678 auto sizeAttr = dyn_cast_or_null<DenseI32ArrayAttr>(
698 Value getBlockSizeZ,
Value dynamicSharedMemorySize,
706 if (!workgroupAttributions.empty())
708 getWorkgroupAttributionsAttrName(
result.name),
712 result.addOperands(asyncDependencies);
717 result.addOperands({gridSizeX, gridSizeY, gridSizeZ, getBlockSizeX,
718 getBlockSizeY, getBlockSizeZ});
720 result.addOperands(clusterSizeX);
722 result.addOperands(clusterSizeY);
724 result.addOperands(clusterSizeZ);
725 if (dynamicSharedMemorySize)
726 result.addOperands(dynamicSharedMemorySize);
728 result.addOperands(asyncObject);
732 result.addAttribute(getModuleAttrName(
result.name), module);
734 result.addAttribute(getFunctionAttrName(
result.name), function);
742 for (
unsigned i = 0; i < kNumConfigRegionAttributes; ++i)
745 for (
Type argTy : workgroupAttributions)
747 for (
Type argTy : privateAttributions)
751 segmentSizes.front() = asyncDependencies.size();
752 segmentSizes[7] = clusterSizeX ? 1 : 0;
753 segmentSizes[8] = clusterSizeY ? 1 : 0;
754 segmentSizes[9] = clusterSizeZ ? 1 : 0;
755 segmentSizes[10] = dynamicSharedMemorySize ? 1 : 0;
756 segmentSizes[11] = asyncObject ? 1 : 0;
757 result.addAttribute(getOperandSegmentSizeAttr(),
762 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
763 auto args = getBody().getArguments();
768 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
769 auto args = getBody().getArguments();
774 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
775 auto args = getBody().getArguments();
780 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
781 auto args = getBody().getArguments();
782 return KernelDim3{args[9], args[10], args[11]};
785std::optional<KernelDim3> LaunchOp::getClusterIds() {
786 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
787 if (!hasClusterSize())
789 auto args = getBody().getArguments();
790 return KernelDim3{args[12], args[13], args[14]};
793std::optional<KernelDim3> LaunchOp::getClusterSize() {
794 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
795 if (!hasClusterSize())
797 auto args = getBody().getArguments();
798 return KernelDim3{args[15], args[16], args[17]};
801KernelDim3 LaunchOp::getGridSizeOperandValues() {
802 auto operands = getOperands().drop_front(getAsyncDependencies().size());
803 return KernelDim3{operands[0], operands[1], operands[2]};
806KernelDim3 LaunchOp::getBlockSizeOperandValues() {
807 auto operands = getOperands().drop_front(getAsyncDependencies().size());
808 return KernelDim3{operands[3], operands[4], operands[5]};
811std::optional<KernelDim3> LaunchOp::getClusterSizeOperandValues() {
812 auto operands = getOperands().drop_front(getAsyncDependencies().size());
813 if (!hasClusterSize())
815 return KernelDim3{operands[6], operands[7], operands[8]};
818template <
typename OpTy>
820 if (!op.getAsyncDependencies().empty() && !op.getAsyncToken())
821 return op.emitOpError(
"dependency operands require the dependency-based "
822 "async model i.e. returning a token");
823 if (op.getAsyncToken() && op.getAsyncObject())
824 return op.emitOpError(
"stream-based and dependency-based async models are "
825 "mutually exclusive");
826 if (op.getNumResults() == 0 && op.getAsyncToken())
827 return op.emitOpError(
"needs to be named when async keyword is specified");
831LogicalResult LaunchOp::verify() {
835 if (!(hasClusterSize()) &&
836 (getClusterSizeX() || getClusterSizeY() || getClusterSizeZ()))
837 return emitOpError() <<
"cluster size must be all present";
841LogicalResult LaunchOp::verifyRegions() {
845 if (getBody().empty()) {
846 return emitOpError(
"body region is empty");
848 unsigned actualNumRegionArgs = getBody().getNumArguments();
849 unsigned expectedNumRegionArgs =
850 getNumConfigRegionAttributes() + getNumWorkgroupAttributions();
851 if (actualNumRegionArgs < expectedNumRegionArgs) {
852 return emitOpError(
"expected at least ")
853 << expectedNumRegionArgs <<
" region arguments, but got "
854 << actualNumRegionArgs;
859 GPUDialect::getWorkgroupAddressSpace())) ||
861 GPUDialect::getPrivateAddressSpace())))
866 for (
Block &block : getBody()) {
869 if (block.back().getNumSuccessors() != 0)
871 if (!isa<gpu::TerminatorOp>(&block.back())) {
874 .append(
"expected '", gpu::TerminatorOp::getOperationName(),
875 "' or a terminator with successors")
876 .attachNote(getLoc())
877 .append(
"in '", LaunchOp::getOperationName(),
"' body region");
890 p <<
'(' << ids.
x <<
", " << ids.
y <<
", " << ids.
z <<
") in (";
891 p << size.
x <<
" = " << operands.
x <<
", ";
892 p << size.
y <<
" = " << operands.
y <<
", ";
893 p << size.
z <<
" = " << operands.
z <<
')';
896void LaunchOp::print(OpAsmPrinter &p) {
897 if (
auto asyncObject = getAsyncObject()) {
898 p <<
" <" << asyncObject <<
" : " << asyncObject.
getType() <<
">";
900 if (getAsyncToken()) {
902 if (!getAsyncDependencies().empty())
903 p <<
" [" << getAsyncDependencies() <<
']';
906 if (hasClusterSize()) {
907 p <<
' ' << getClustersKeyword();
909 getClusterSizeOperandValues().value(),
910 getClusterIds().value());
912 p <<
' ' << getBlocksKeyword();
915 p <<
' ' << getThreadsKeyword();
918 if (getDynamicSharedMemorySize())
919 p <<
' ' << getDynamicSharedMemorySizeKeyword() <<
' '
920 << getDynamicSharedMemorySize();
923 StringRef moduleAttrName = getModuleAttrName();
924 if (
auto module = getModule()) {
925 p <<
' ' << moduleAttrName <<
'(';
930 StringRef functionAttrName = getFunctionAttrName();
931 if (
auto function = getFunction()) {
932 p <<
' ' << functionAttrName <<
'(';
937 if (getCooperative())
947 (*this)->getDiscardableAttrDictionary().getValue(), {
948 LaunchOp::getOperandSegmentSizeAttr(),
949 getWorkgroupAttributionsAttrName(), getCooperativeAttrName(),
950 moduleAttrName, functionAttrName});
965 assert(
indices.size() == 3 &&
"space for three indices expected");
972 if (args.size() != 3) {
974 << keyword <<
" expects 3 arguments, but got " << args.size();
976 std::move(args.begin(), args.end(),
indices.begin());
978 for (
int i = 0; i < 3; ++i) {
1001ParseResult LaunchOp::parse(OpAsmParser &parser, OperationState &
result) {
1003 SmallVector<OpAsmParser::UnresolvedOperand, LaunchOp::kNumConfigOperands>
1004 sizes(LaunchOp::kNumConfigOperands);
1007 SmallVector<OpAsmParser::UnresolvedOperand, 16> regionArgs(
1008 LaunchOp::kNumConfigRegionAttributes);
1011 OpAsmParser::UnresolvedOperand asyncObjectOperand;
1012 Type asyncObjectType;
1013 bool hasAsyncObject =
false;
1015 hasAsyncObject =
true;
1022 SmallVector<OpAsmParser::UnresolvedOperand, 4> asyncDependencies;
1023 Type asyncTokenType;
1030 if (!asyncTokenType)
1033 "gpu.launch requires 'async' keyword to return a value");
1034 result.types.push_back(asyncTokenType);
1037 bool hasCluster =
false;
1041 regionArgs.resize(18);
1043 MutableArrayRef<OpAsmParser::UnresolvedOperand> sizesRef(sizes);
1044 MutableArrayRef<OpAsmParser::UnresolvedOperand> regionArgsRef(regionArgs);
1050 parser, sizesRef.drop_front(6), regionArgsRef.slice(15, 3),
1051 regionArgsRef.slice(12, 3), LaunchOp::getClustersKeyword()))
1059 if (parser.
parseKeyword(LaunchOp::getBlocksKeyword()) ||
1061 regionArgsRef.slice(6, 3), regionArgsRef.slice(0, 3),
1062 LaunchOp::getBlocksKeyword()) ||
1065 regionArgsRef.slice(9, 3), regionArgsRef.slice(3, 3),
1066 LaunchOp::getThreadsKeyword()) ||
1071 OpAsmParser::UnresolvedOperand dynamicSharedMemorySize;
1072 bool hasDynamicSharedMemorySize =
false;
1074 LaunchOp::getDynamicSharedMemorySizeKeyword())) {
1075 hasDynamicSharedMemorySize =
true;
1085 asyncObjectType,
result.operands))
1089 StringRef moduleAttrName = getModuleAttrName(
result.name);
1091 FlatSymbolRefAttr moduleSymbol;
1099 StringRef functionAttrName = getFunctionAttrName(
result.name);
1101 FlatSymbolRefAttr funcSymbol;
1117 SmallVector<Type, LaunchOp::kNumConfigRegionAttributes> dataTypes(
1118 LaunchOp::kNumConfigRegionAttributes + 6, index);
1120 SmallVector<OpAsmParser::Argument> regionArguments;
1121 for (
auto ssaValueAndType : llvm::zip(regionArgs, dataTypes)) {
1122 OpAsmParser::Argument arg;
1123 arg.
ssaName = std::get<0>(ssaValueAndType);
1124 arg.
type = std::get<1>(ssaValueAndType);
1125 regionArguments.push_back(arg);
1136 unsigned numWorkgroupAttrs = regionArguments.size() -
1137 LaunchOp::kNumConfigRegionAttributes -
1138 (hasCluster ? 6 : 0);
1139 if (numWorkgroupAttrs != 0)
1140 result.addAttribute(LaunchOp::getWorkgroupAttributionsAttrName(
result.name),
1151 Region *body =
result.addRegion();
1156 SmallVector<int32_t, 12> segmentSizes(12, 1);
1157 segmentSizes.front() = asyncDependencies.size();
1160 segmentSizes[7] = 0;
1161 segmentSizes[8] = 0;
1162 segmentSizes[9] = 0;
1164 segmentSizes[10] = hasDynamicSharedMemorySize ? 1 : 0;
1165 segmentSizes[11] = hasAsyncObject ? 1 : 0;
1166 result.addAttribute(LaunchOp::getOperandSegmentSizeAttr(),
1180 bool simplified =
false;
1181 auto constPropIdUses = [&](
Value id,
Value size) {
1185 if (
id.getUses().empty())
1197 constPropIdUses(op.getBlockIds().x, op.getGridSizeX());
1198 constPropIdUses(op.getBlockIds().y, op.getGridSizeY());
1199 constPropIdUses(op.getBlockIds().z, op.getGridSizeZ());
1200 constPropIdUses(op.getThreadIds().x, op.getBlockSizeX());
1201 constPropIdUses(op.getThreadIds().y, op.getBlockSizeY());
1202 constPropIdUses(op.getThreadIds().z, op.getBlockSizeZ());
1208void LaunchOp::getCanonicalizationPatterns(RewritePatternSet &rewrites,
1209 MLIRContext *context) {
1210 rewrites.
add<FoldLaunchArguments>(context);
1215BlockArgument LaunchOp::addWorkgroupAttribution(Type type, Location loc) {
1216 int64_t cur = getWorkgroupAttributions().value_or(0);
1217 setWorkgroupAttributions(std::optional<int64_t>(cur + 1));
1218 return getBody().insertArgument(
1219 getNumConfigRegionAttributes() +
static_cast<unsigned>(cur), type, loc);
1224BlockArgument LaunchOp::addPrivateAttribution(Type type, Location loc) {
1227 return getBody().addArgument(type, loc);
1234void LaunchFuncOp::build(OpBuilder &builder, OperationState &
result,
1235 SymbolRefAttr kernelSymbol,
KernelDim3 gridSize,
1236 KernelDim3 getBlockSize, Value dynamicSharedMemorySize,
1237 ValueRange kernelOperands, Type asyncTokenType,
1238 ValueRange asyncDependencies, Value asyncObject,
1239 std::optional<KernelDim3> clusterSize) {
1240 assert(kernelSymbol.getNestedReferences().size() == 1 &&
1241 "expected a symbol reference with a single nested reference");
1242 result.addOperands(asyncDependencies);
1249 if (clusterSize.has_value())
1250 result.addOperands({clusterSize->x, clusterSize->y, clusterSize->z});
1251 if (dynamicSharedMemorySize)
1252 result.addOperands(dynamicSharedMemorySize);
1253 result.addOperands(kernelOperands);
1255 result.addOperands(asyncObject);
1257 Properties &prop =
result.getOrAddProperties<Properties>();
1258 prop.kernel = kernelSymbol;
1259 size_t segmentSizesLen = std::size(prop.operandSegmentSizes);
1261 llvm::fill(prop.operandSegmentSizes, 1);
1262 prop.operandSegmentSizes[0] = asyncDependencies.size();
1263 if (!clusterSize.has_value()) {
1264 prop.operandSegmentSizes[segmentSizesLen - 4] = 0;
1265 prop.operandSegmentSizes[segmentSizesLen - 5] = 0;
1266 prop.operandSegmentSizes[segmentSizesLen - 6] = 0;
1268 prop.operandSegmentSizes[segmentSizesLen - 3] =
1269 dynamicSharedMemorySize ? 1 : 0;
1270 prop.operandSegmentSizes[segmentSizesLen - 2] =
1271 static_cast<int32_t
>(kernelOperands.size());
1272 prop.operandSegmentSizes[segmentSizesLen - 1] = asyncObject ? 1 : 0;
1275void LaunchFuncOp::build(OpBuilder &builder, OperationState &
result,
1277 KernelDim3 getBlockSize, Value dynamicSharedMemorySize,
1278 ValueRange kernelOperands, Type asyncTokenType,
1279 ValueRange asyncDependencies, Value asyncObject,
1280 std::optional<KernelDim3> clusterSize) {
1281 auto kernelModule = kernelFunc->getParentOfType<GPUModuleOp>();
1283 SymbolRefAttr::get(kernelModule.getNameAttr(),
1284 {SymbolRefAttr::get(kernelFunc.getNameAttr())});
1285 build(builder,
result, kernelSymbol, gridSize, getBlockSize,
1286 dynamicSharedMemorySize, kernelOperands, asyncTokenType,
1287 asyncDependencies, asyncObject, clusterSize);
1290StringAttr LaunchFuncOp::getKernelModuleName() {
1294StringAttr LaunchFuncOp::getKernelName() {
1298unsigned LaunchFuncOp::getNumKernelOperands() {
1299 return getKernelOperands().size();
1302Value LaunchFuncOp::getKernelOperand(
unsigned i) {
1303 return getKernelOperands()[i];
1306KernelDim3 LaunchFuncOp::getGridSizeOperandValues() {
1307 auto operands = getOperands().drop_front(getAsyncDependencies().size());
1308 return KernelDim3{operands[0], operands[1], operands[2]};
1311KernelDim3 LaunchFuncOp::getBlockSizeOperandValues() {
1312 auto operands = getOperands().drop_front(getAsyncDependencies().size());
1313 return KernelDim3{operands[3], operands[4], operands[5]};
1316KernelDim3 LaunchFuncOp::getClusterSizeOperandValues() {
1317 assert(hasClusterSize() &&
1318 "cluster size is not set, check hasClusterSize() first");
1319 auto operands = getOperands().drop_front(getAsyncDependencies().size());
1320 return KernelDim3{operands[6], operands[7], operands[8]};
1323LogicalResult LaunchFuncOp::verify() {
1327 auto module = (*this)->getParentOfType<ModuleOp>();
1329 return emitOpError(
"expected to belong to a module");
1331 if (!module->getDiscardableAttrOfType<UnitAttr>(
1332 GPUDialect::getContainerModuleAttrName()))
1333 return emitOpError(
"expected the closest surrounding module to have the '" +
1334 GPUDialect::getContainerModuleAttrName() +
1337 if (hasClusterSize()) {
1338 if (getClusterSizeY().
getType() != getClusterSizeX().
getType() ||
1340 return emitOpError()
1341 <<
"expects types of the cluster dimensions must be the same";
1348LaunchFuncOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
1349 LaunchFuncOp launchOp = *
this;
1352 if (isa<GPUModuleOp>(table))
1357 if (!launchOp->getParentOp() ||
1358 launchOp->getParentOp()->getParentOp() != table)
1363 if (!launchOp.getKernelAttr())
1367 StringAttr kernelContainerName = launchOp.getKernelModuleName();
1368 Operation *kernelContainer =
1370 if (!kernelContainer)
1372 <<
"kernel container '" << kernelContainerName.getValue()
1373 <<
"' is undefined";
1376 if (isa<BinaryOp>(kernelContainer))
1379 auto kernelModule = dyn_cast<GPUModuleOp>(kernelContainer);
1381 return launchOp.emitOpError()
1382 <<
"kernel module '" << kernelContainerName.getValue()
1383 <<
"' is undefined";
1387 kernelModule, launchOp.getKernelName());
1390 << launchOp.getKernel() <<
"' is undefined";
1391 auto kernelConvertedFunction = dyn_cast<FunctionOpInterface>(kernelFunc);
1392 if (!kernelConvertedFunction) {
1393 InFlightDiagnostic
diag = launchOp.emitOpError()
1394 <<
"referenced kernel '" << launchOp.getKernel()
1395 <<
"' is not a function";
1396 diag.attachNote(kernelFunc->
getLoc()) <<
"see the kernel definition here";
1400 if (!GPUDialect::isKernel(kernelFunc))
1401 return launchOp.emitOpError(
"kernel function is missing the '")
1402 << GPUDialect::getKernelFuncAttrName() <<
"' attribute";
1407 auto kernelGPUFunction = dyn_cast<gpu::GPUFuncOp>(kernelFunc);
1408 if (!kernelGPUFunction)
1411 unsigned actualNumArguments = launchOp.getNumKernelOperands();
1412 unsigned expectedNumArguments = kernelGPUFunction.getNumArguments();
1413 if (expectedNumArguments != actualNumArguments)
1414 return launchOp.emitOpError(
"got ")
1415 << actualNumArguments <<
" kernel operands but expected "
1416 << expectedNumArguments;
1418 FunctionType functionType = kernelGPUFunction.getFunctionType();
1419 for (
unsigned i = 0; i < expectedNumArguments; ++i) {
1420 if (launchOp.getKernelOperand(i).getType() != functionType.getInput(i)) {
1421 return launchOp.emitOpError(
"type of function argument ")
1422 << i <<
" does not match";
1431 std::optional<OpAsmParser::UnresolvedOperand> clusterValue,
1432 Type &clusterXTy,
Type &clusterYTy,
Type &clusterZTy) {
1439 if (clusterValue.has_value()) {
1440 clusterXTy = clusterYTy = clusterZTy = dimTy;
1447 Type clusterYTy,
Type clusterZTy) {
1449 printer <<
": " << dimTy;
1459 auto parseElement = [&]() -> ParseResult {
1460 return failure(parser.
parseOperand(argNames.emplace_back()) ||
1465 parseElement,
" in argument list");
1470 if (operands.empty())
1473 llvm::interleaveComma(llvm::zip_equal(operands, types), printer,
1474 [&](
const auto &pair) {
1475 auto [operand, type] = pair;
1476 printer << operand <<
" : " << type;
1485void ShuffleOp::build(OpBuilder &builder, OperationState &
result, Value value,
1486 int32_t offset, int32_t width, ShuffleMode mode) {
1487 build(builder,
result, value,
1488 arith::ConstantOp::create(builder,
result.location,
1490 arith::ConstantOp::create(builder,
result.location,
1499LogicalResult RotateOp::verify() {
1500 uint32_t offset = getOffset();
1501 uint32_t width = getWidth();
1503 if (offset >= width) {
1504 return emitOpError() <<
"offset must be in the range [0, " << width <<
")";
1514LogicalResult BarrierOp::verify() {
1515 BarrierScope scope = getScope();
1517 if (getNamedBarrier() && scope != BarrierScope::Workgroup)
1518 return emitOpError(
"named barriers require workgroup scope");
1526 auto nextOp = dyn_cast_or_null<BarrierOp>(op->getNextNode());
1531 if (op.getScope() != nextOp.getScope())
1535 if (op.getNamedBarrier() != nextOp.getNamedBarrier())
1538 std::optional<ArrayAttr> thisMemfence = op.getAddressSpaces();
1539 std::optional<ArrayAttr> nextMemfence = nextOp.getAddressSpaces();
1543 if (!nextMemfence) {
1544 op.removeAddressSpacesAttr();
1548 if (*thisMemfence == *nextMemfence) {
1552 llvm::SmallSetVector<Attribute, 4> mergedSpaces;
1554 mergedSpaces.insert(attr);
1556 mergedSpaces.insert(attr);
1557 op.setAddressSpacesAttr(rewriter.
getArrayAttr(mergedSpaces.takeVector()));
1565void BarrierOp::getCanonicalizationPatterns(RewritePatternSet &results,
1566 MLIRContext *context) {
1570void BarrierOp::build(mlir::OpBuilder &odsBuilder,
1571 mlir::OperationState &odsState,
1572 std::optional<AddressSpace> addressSpace) {
1576 AddressSpaceAttr::get(odsBuilder.
getContext(), addressSpace.value()));
1578 odsBuilder, odsState, addressSpacesAttr, Value{},
1579 BarrierScopeAttr::get(odsBuilder.
getContext(), BarrierScope::Workgroup));
1586void BarrierOp::build(OpBuilder &builder, OperationState &odsState,
1587 Value memrefToFence) {
1588 std::optional<AddressSpace> addrSpaceToFence;
1589 if (
auto memrefType = dyn_cast<BaseMemRefType>(memrefToFence.
getType()))
1590 if (
auto addrSpaceAttr = dyn_cast_if_present<gpu::AddressSpaceAttr>(
1591 memrefType.getMemorySpace()))
1592 addrSpaceToFence = addrSpaceAttr.getValue();
1593 return build(builder, odsState, addrSpaceToFence);
1602BlockArgument GPUFuncOp::addWorkgroupAttribution(Type type, Location loc) {
1603 int64_t cur = getWorkgroupAttributions().value_or(0);
1604 setWorkgroupAttributions(std::optional<int64_t>(cur + 1));
1605 return getBody().insertArgument(
1606 getFunctionType().getNumInputs() +
static_cast<unsigned>(cur), type, loc);
1611BlockArgument GPUFuncOp::addPrivateAttribution(Type type, Location loc) {
1614 return getBody().addArgument(type, loc);
1617void GPUFuncOp::build(OpBuilder &builder, OperationState &
result,
1618 StringRef name, FunctionType type,
1621 ArrayRef<NamedAttribute> attrs) {
1622 OpBuilder::InsertionGuard g(builder);
1624 result.getOrAddProperties<Properties>().sym_name =
1626 result.addAttribute(getFunctionTypeAttrName(
result.name),
1627 TypeAttr::get(type));
1628 result.addAttribute(getWorkgroupAttributionsAttrName(
result.name),
1630 result.addAttributes(attrs);
1631 Region *body =
result.addRegion();
1635 for (Type argTy : type.getInputs())
1637 for (Type argTy : workgroupAttributions)
1639 for (Type argTy : privateAttributions)
1658 size_t existingArgs = args.size();
1665 bool hadAttrs = llvm::any_of(
ArrayRef(args).drop_front(existingArgs),
1670 attributionAttrs =
nullptr;
1676 for (
const auto &argument :
ArrayRef(args).drop_front(existingArgs)) {
1677 if (!argument.attrs)
1680 attributionAttrsVec.push_back(argument.attrs);
1682 attributionAttrs = builder.
getArrayAttr(attributionAttrsVec);
1691ParseResult GPUFuncOp::parse(OpAsmParser &parser, OperationState &
result) {
1692 SmallVector<OpAsmParser::Argument> entryArgs;
1693 SmallVector<DictionaryAttr> resultAttrs;
1694 SmallVector<Type> resultTypes;
1698 StringAttr nameAttr;
1705 parser,
false, entryArgs, isVariadic, resultTypes,
1709 if (!entryArgs.empty() && entryArgs[0].ssaName.name.empty())
1710 return parser.
emitError(signatureLocation)
1711 <<
"gpu.func requires named arguments";
1717 SmallVector<Type> argTypes;
1718 for (
auto &arg : entryArgs)
1719 argTypes.push_back(arg.
type);
1721 result.addAttribute(getFunctionTypeAttrName(
result.name),
1722 TypeAttr::get(type));
1725 builder,
result, entryArgs, resultAttrs, getArgAttrsAttrName(
result.name),
1726 getResAttrsAttrName(
result.name));
1728 Attribute workgroupAttributionAttrs;
1731 entryArgs, workgroupAttributionAttrs)))
1736 unsigned numWorkgroupAttrs = entryArgs.size() - type.getNumInputs();
1737 if (numWorkgroupAttrs != 0)
1739 GPUFuncOp::getWorkgroupAttributionsAttrName(
result.name),
1741 if (workgroupAttributionAttrs)
1742 result.addAttribute(GPUFuncOp::getWorkgroupAttribAttrsAttrName(
result.name),
1743 workgroupAttributionAttrs);
1745 Attribute privateAttributionAttrs;
1748 entryArgs, privateAttributionAttrs)))
1750 if (privateAttributionAttrs)
1751 result.addAttribute(GPUFuncOp::getPrivateAttribAttrsAttrName(
result.name),
1752 privateAttributionAttrs);
1756 result.addAttribute(GPUFuncOp::getKernelAttrName(
result.name),
1765 auto *body =
result.addRegion();
1769void GPUFuncOp::print(OpAsmPrinter &p) {
1773 FunctionType type = getFunctionType();
1779 getWorkgroupAttribAttrs().value_or(
nullptr));
1781 getPrivateAttribAttrs().value_or(
nullptr));
1783 p <<
' ' << getKernelKeyword();
1787 {getWorkgroupAttributionsAttrName(), getKernelAttrName(),
1788 GPUDialect::getKernelFuncAttrName(), getFunctionTypeAttrName(),
1789 getArgAttrsAttrName(), getResAttrsAttrName(),
1790 getWorkgroupAttribAttrsAttrName(), getPrivateAttribAttrsAttrName()});
1796 StringAttr attrName) {
1797 ArrayAttr allAttrs = attrName == op.getWorkgroupAttribAttrsAttrName()
1798 ? op.getWorkgroupAttribAttrsAttr()
1799 : op.getPrivateAttribAttrsAttr();
1800 if (!allAttrs ||
index >= allAttrs.size())
1801 return DictionaryAttr();
1802 return llvm::cast<DictionaryAttr>(allAttrs[
index]);
1805DictionaryAttr GPUFuncOp::getworkgroupAttributionAttrs(
unsigned index) {
1809DictionaryAttr GPUFuncOp::getPrivateAttributionAttrs(
unsigned index) {
1814 DictionaryAttr value, StringAttr attrName) {
1816 ArrayAttr allAttrs = attrName == op.getWorkgroupAttribAttrsAttrName()
1817 ? op.getWorkgroupAttribAttrsAttr()
1818 : op.getPrivateAttribAttrsAttr();
1821 elements.append(allAttrs.begin(), allAttrs.end());
1822 while (elements.size() <=
index)
1823 elements.push_back(DictionaryAttr::get(ctx));
1825 elements[
index] = DictionaryAttr::get(ctx);
1827 elements[
index] = value;
1828 ArrayAttr newValue = ArrayAttr::get(ctx, elements);
1829 if (attrName == op.getWorkgroupAttribAttrsAttrName())
1830 op.setWorkgroupAttribAttrsAttr(newValue);
1832 op.setPrivateAttribAttrsAttr(newValue);
1835void GPUFuncOp::setworkgroupAttributionAttrs(
unsigned index,
1836 DictionaryAttr value) {
1840void GPUFuncOp::setPrivateAttributionAttrs(
unsigned int index,
1841 DictionaryAttr value) {
1846 StringAttr name, StringAttr attrsName) {
1850 return dict.get(name);
1853Attribute GPUFuncOp::getWorkgroupAttributionAttr(
unsigned index,
1855 assert(index < getNumWorkgroupAttributions() &&
1856 "index must map to a workgroup attribution");
1858 getWorkgroupAttribAttrsAttrName());
1861Attribute GPUFuncOp::getPrivateAttributionAttr(
unsigned index,
1863 assert(index < getNumPrivateAttributions() &&
1864 "index must map to a private attribution");
1866 getPrivateAttribAttrsAttrName());
1870 Attribute value, StringAttr attrsName) {
1875 elems.append(oldDict.getValue().begin(), oldDict.getValue().end());
1878 bool mustSort =
true;
1879 for (
unsigned i = 0, e = elems.size(); i < e; ++i) {
1880 if (elems[i].getName() == name) {
1883 std::swap(elems[i], elems[elems.size() - 1]);
1895 elems.emplace_back(name, value);
1898 DictionaryAttr::sortInPlace(elems);
1900 auto newDict = DictionaryAttr::getWithSorted(ctx, elems);
1904void GPUFuncOp::setWorkgroupAttributionAttr(
unsigned index, StringAttr name,
1906 assert(index < getNumWorkgroupAttributions() &&
1907 "index must map to a workgroup attribution");
1909 getWorkgroupAttribAttrsAttrName());
1912void GPUFuncOp::setPrivateAttributionAttr(
unsigned index, StringAttr name,
1914 assert(index < getNumPrivateAttributions() &&
1915 "index must map to a private attribution");
1917 getPrivateAttribAttrsAttrName());
1920LogicalResult GPUFuncOp::verifyType() {
1921 if (isKernel() && getFunctionType().getNumResults() != 0)
1922 return emitOpError() <<
"expected void return type for kernel function";
1928LogicalResult GPUFuncOp::verifyBody() {
1930 return emitOpError() <<
"expected body with at least one block";
1931 unsigned numFuncArguments = getNumArguments();
1932 unsigned numWorkgroupAttributions = getNumWorkgroupAttributions();
1933 unsigned numBlockArguments = front().getNumArguments();
1934 if (numBlockArguments < numFuncArguments + numWorkgroupAttributions)
1935 return emitOpError() <<
"expected at least "
1936 << numFuncArguments + numWorkgroupAttributions
1937 <<
" arguments to body region";
1939 ArrayRef<Type> funcArgTypes = getFunctionType().getInputs();
1940 for (
unsigned i = 0; i < numFuncArguments; ++i) {
1941 Type blockArgType = front().getArgument(i).getType();
1942 if (funcArgTypes[i] != blockArgType)
1943 return emitOpError() <<
"expected body region argument #" << i
1944 <<
" to be of type " << funcArgTypes[i] <<
", got "
1949 GPUDialect::getWorkgroupAddressSpace())) ||
1951 GPUDialect::getPrivateAddressSpace())))
1961LogicalResult gpu::ReturnOp::verify() {
1962 GPUFuncOp function = (*this)->getParentOfType<GPUFuncOp>();
1964 FunctionType funType = function.getFunctionType();
1966 if (funType.getNumResults() != getOperands().size())
1967 return emitOpError()
1968 .append(
"expected ", funType.getNumResults(),
" result operands")
1969 .attachNote(function.getLoc())
1970 .append(
"return type declared here");
1972 for (
const auto &pair : llvm::enumerate(
1973 llvm::zip(function.getFunctionType().getResults(), getOperands()))) {
1974 auto [type, operand] = pair.value();
1975 if (type != operand.getType())
1976 return emitOpError() <<
"unexpected type `" << operand.getType()
1977 <<
"' for operand #" << pair.index();
1986void GPUModuleOp::build(OpBuilder &builder, OperationState &
result,
1988 Attribute offloadingHandler) {
1989 result.addRegion()->emplaceBlock();
1990 Properties &props =
result.getOrAddProperties<Properties>();
1992 props.targets = targets;
1994 props.offloadingHandler = offloadingHandler;
1997void GPUModuleOp::build(OpBuilder &builder, OperationState &
result,
1998 StringRef name, ArrayRef<Attribute> targets,
1999 Attribute offloadingHandler) {
2000 build(builder,
result, name,
2005bool GPUModuleOp::hasTarget(Attribute
target) {
2006 if (
ArrayAttr targets = getTargetsAttr())
2007 return llvm::count(targets.getValue(),
target);
2011void GPUModuleOp::setTargets(ArrayRef<TargetAttrInterface> targets) {
2012 ArrayAttr &targetsAttr = getProperties().targets;
2013 SmallVector<Attribute> targetsVector(targets);
2014 targetsAttr = ArrayAttr::get(
getContext(), targetsVector);
2017LogicalResult GPUModuleOp::verify() {
2018 auto targets = getTargetsAttr();
2023 for (
auto target : targets) {
2024 if (
auto verifyTargetAttr =
2025 llvm::dyn_cast<TargetAttrVerifyInterface>(
target)) {
2026 if (verifyTargetAttr.verifyTarget(getOperation()).
failed())
2036void BinaryOp::build(OpBuilder &builder, OperationState &
result, StringRef name,
2037 Attribute offloadingHandler,
ArrayAttr objects) {
2038 auto &properties =
result.getOrAddProperties<Properties>();
2041 properties.objects = objects;
2042 if (offloadingHandler)
2043 properties.offloadingHandler = offloadingHandler;
2045 properties.offloadingHandler = builder.
getAttr<SelectObjectAttr>(
nullptr);
2048void BinaryOp::build(OpBuilder &builder, OperationState &
result, StringRef name,
2049 Attribute offloadingHandler, ArrayRef<Attribute> objects) {
2050 build(builder,
result, name, offloadingHandler,
2062 if (!offloadingHandler)
2069 if (offloadingHandler != SelectObjectAttr::get(op->
getContext(),
nullptr))
2070 printer <<
'<' << offloadingHandler <<
'>';
2077LogicalResult MemcpyOp::verify() {
2078 auto srcType = getSrc().getType();
2079 auto dstType = getDst().getType();
2082 return emitOpError(
"arguments have incompatible element type");
2085 return emitOpError(
"arguments have incompatible shape");
2094struct EraseTrivialCopyOp :
public OpRewritePattern<MemcpyOp> {
2095 using OpRewritePattern<MemcpyOp>::OpRewritePattern;
2097 LogicalResult matchAndRewrite(MemcpyOp op,
2098 PatternRewriter &rewriter)
const override {
2099 Value dest = op.getDst();
2108 if (llvm::any_of(dest.
getUsers(), [op, dest](Operation *user) {
2109 return user != op &&
2110 !hasSingleEffect<MemoryEffects::Free>(user, dest);
2116 if (op.getAsyncDependencies().size() > 1 ||
2117 ((op.getAsyncDependencies().empty() && op.getAsyncToken()) ||
2118 (!op.getAsyncDependencies().empty() && !op.getAsyncToken())))
2120 rewriter.
replaceOp(op, op.getAsyncDependencies());
2127void MemcpyOp::getCanonicalizationPatterns(RewritePatternSet &results,
2128 MLIRContext *context) {
2129 results.
add<EraseTrivialCopyOp>(context);
2136LogicalResult SubgroupMmaLoadMatrixOp::verify() {
2137 auto srcType = getSrcMemref().getType();
2138 auto resType = getRes().getType();
2139 auto resMatrixType = llvm::cast<gpu::MMAMatrixType>(resType);
2140 auto operand = resMatrixType.getOperand();
2141 auto srcMemrefType = llvm::cast<MemRefType>(srcType);
2143 if (!srcMemrefType.isLastDimUnitStride())
2145 "expected source memref most minor dim must have unit stride");
2147 if (operand !=
"AOp" && operand !=
"BOp" && operand !=
"COp")
2148 return emitError(
"only AOp, BOp and COp can be loaded");
2157LogicalResult SubgroupMmaStoreMatrixOp::verify() {
2158 auto srcType = getSrc().getType();
2159 auto dstType = getDstMemref().getType();
2160 auto srcMatrixType = llvm::cast<gpu::MMAMatrixType>(srcType);
2161 auto dstMemrefType = llvm::cast<MemRefType>(dstType);
2163 if (!dstMemrefType.isLastDimUnitStride())
2165 "expected destination memref most minor dim must have unit stride");
2167 if (srcMatrixType.getOperand() !=
"COp")
2169 "expected the operand matrix being stored to have 'COp' operand type");
2178LogicalResult SubgroupMmaComputeOp::verify() {
2179 enum OperandMap {
A,
B,
C };
2180 SmallVector<MMAMatrixType, 3> opTypes;
2181 opTypes.push_back(llvm::cast<MMAMatrixType>(getOpA().
getType()));
2182 opTypes.push_back(llvm::cast<MMAMatrixType>(getOpB().
getType()));
2183 opTypes.push_back(llvm::cast<MMAMatrixType>(getOpC().
getType()));
2185 if (opTypes[A].getOperand() !=
"AOp" || opTypes[B].getOperand() !=
"BOp" ||
2186 opTypes[C].getOperand() !=
"COp")
2187 return emitError(
"operands must be in the order AOp, BOp, COp");
2189 ArrayRef<int64_t> aShape, bShape, cShape;
2190 aShape = opTypes[
A].getShape();
2191 bShape = opTypes[
B].getShape();
2192 cShape = opTypes[
C].getShape();
2194 if (aShape[1] != bShape[0] || aShape[0] != cShape[0] ||
2195 bShape[1] != cShape[1])
2196 return emitError(
"operand shapes do not satisfy matmul constraints");
2201LogicalResult MemcpyOp::fold(FoldAdaptor adaptor,
2202 SmallVectorImpl<::mlir::OpFoldResult> &results) {
2206LogicalResult MemsetOp::fold(FoldAdaptor adaptor,
2207 SmallVectorImpl<::mlir::OpFoldResult> &results) {
2220struct EraseRedundantGpuWaitOpPairs :
public OpRewritePattern<WaitOp> {
2224 LogicalResult matchAndRewrite(WaitOp op,
2225 PatternRewriter &rewriter)
const final {
2226 auto predicate = [](Value value) {
2227 auto waitOp = value.getDefiningOp<WaitOp>();
2228 return waitOp && waitOp->getNumOperands() == 0;
2230 if (llvm::none_of(op.getAsyncDependencies(), predicate))
2232 SmallVector<Value> validOperands;
2233 for (Value operand : op->getOperands()) {
2234 if (predicate(operand))
2236 validOperands.push_back(operand);
2238 rewriter.
modifyOpInPlace(op, [&]() { op->setOperands(validOperands); });
2250struct SimplifyGpuWaitOp :
public OpRewritePattern<WaitOp> {
2254 LogicalResult matchAndRewrite(WaitOp op,
2255 PatternRewriter &rewriter)
const final {
2258 if (op.getAsyncDependencies().empty() && !op.getAsyncToken()) {
2263 if (llvm::hasSingleElement(op.getAsyncDependencies()) &&
2264 op.getAsyncToken()) {
2265 rewriter.
replaceOp(op, op.getAsyncDependencies());
2269 if (op.getAsyncToken() && op.getAsyncToken().use_empty()) {
2279void WaitOp::getCanonicalizationPatterns(RewritePatternSet &results,
2280 MLIRContext *context) {
2281 results.
add<EraseRedundantGpuWaitOpPairs, SimplifyGpuWaitOp>(context);
2288LogicalResult AllocOp::verify() {
2289 auto memRefType = llvm::cast<MemRefType>(getMemref().
getType());
2295 unsigned numSymbols = 0;
2296 if (!memRefType.getLayout().isIdentity())
2297 numSymbols = memRefType.getLayout().getAffineMap().getNumSymbols();
2298 if (getSymbolOperands().size() != numSymbols) {
2300 "symbol operand count does not equal memref symbol count");
2310struct SimplifyDimOfAllocOp :
public OpRewritePattern<memref::DimOp> {
2311 using OpRewritePattern<memref::DimOp>::OpRewritePattern;
2313 LogicalResult matchAndRewrite(memref::DimOp dimOp,
2314 PatternRewriter &rewriter)
const override {
2315 std::optional<int64_t> index = dimOp.getConstantIndex();
2319 int64_t indexVal = index.value();
2320 auto memrefType = llvm::dyn_cast<MemRefType>(dimOp.getSource().getType());
2321 if (!memrefType || indexVal < 0 || indexVal >= memrefType.getRank() ||
2322 !memrefType.isDynamicDim(indexVal))
2325 auto alloc = dimOp.getSource().getDefiningOp<AllocOp>();
2329 Value substituteOp = *(alloc.getDynamicSizes().begin() +
2330 memrefType.getDynamicDimIndex(indexVal));
2331 rewriter.
replaceOp(dimOp, substituteOp);
2338void AllocOp::getCanonicalizationPatterns(RewritePatternSet &results,
2339 MLIRContext *context) {
2340 results.
add<SimplifyDimOfAllocOp>(context);
2348 Attribute
target, CompilationTarget format,
2349 StringAttr
object, DictionaryAttr properties,
2350 KernelTableAttr kernels) {
2352 return emitError() <<
"the target attribute cannot be null";
2353 if (
target.hasPromiseOrImplementsInterface<TargetAttrInterface>())
2355 return emitError() <<
"the target attribute must implement or promise the "
2356 "`gpu::TargetAttrInterface`";
2360ParseResult parseObject(AsmParser &odsParser, CompilationTarget &format,
2361 StringAttr &
object) {
2362 std::optional<CompilationTarget> formatResult;
2363 StringRef enumKeyword;
2366 formatResult = CompilationTarget::Fatbin;
2367 if (!formatResult &&
2369 gpu::symbolizeEnum<gpu::CompilationTarget>(enumKeyword)) &&
2371 return odsParser.
emitError(loc,
"expected an equal sign");
2373 return odsParser.
emitError(loc,
"expected keyword for GPU object format");
2374 FailureOr<StringAttr> objectResult =
2375 FieldParser<StringAttr>::parse(odsParser);
2376 if (
failed(objectResult))
2378 "failed to parse GPU_ObjectAttr parameter "
2379 "'object' which is to be a `StringAttr`");
2380 format = *formatResult;
2381 object = *objectResult;
2385void printObject(AsmPrinter &odsParser, CompilationTarget format,
2386 StringAttr
object) {
2387 if (format != CompilationTarget::Fatbin)
2388 odsParser << stringifyEnum(format) <<
" = ";
2389 odsParser << object;
2402 if (
auto intAttr = mlir::dyn_cast<IntegerAttr>(
target)) {
2403 if (intAttr.getInt() < 0) {
2404 return emitError() <<
"the object index must be positive";
2406 }
else if (!
target.hasPromiseOrImplementsInterface<TargetAttrInterface>()) {
2408 <<
"the target attribute must be a GPU Target attribute";
2418LogicalResult gpu::DynamicSharedMemoryOp::verify() {
2419 if (!getOperation()->getParentWithTrait<OpTrait::SymbolTable>())
2420 return emitOpError() <<
"must be inside an op with symbol table";
2422 MemRefType memrefType = getResultMemref().getType();
2424 if (!GPUDialect::hasWorkgroupMemoryAddressSpace(memrefType)) {
2425 return emitOpError() <<
"address space must be "
2426 << gpu::AddressSpaceAttr::getMnemonic() <<
"<"
2427 << stringifyEnum(gpu::AddressSpace::Workgroup) <<
">";
2429 if (memrefType.hasStaticShape()) {
2430 return emitOpError() <<
"result memref type must be memref<?xi8, "
2431 "#gpu.address_space<workgroup>>";
2440void WarpExecuteOnLane0Op::print(OpAsmPrinter &p) {
2441 p <<
"(" << getLaneid() <<
")";
2443 SmallVector<StringRef> coreAttr = {getWarpSizeAttrName()};
2444 p <<
"[" << getWarpSize() <<
"]";
2446 if (!getArgs().empty())
2447 p <<
" args(" << getArgs() <<
" : " << getArgs().getTypes() <<
")";
2448 if (!getResults().empty())
2449 p <<
" -> (" << getResults().getTypes() <<
')';
2453 !getResults().empty());
2455 getOperation()->getDiscardableAttrDictionary().getValue(), coreAttr);
2458ParseResult WarpExecuteOnLane0Op::parse(OpAsmParser &parser,
2459 OperationState &
result) {
2461 result.regions.reserve(1);
2462 Region *warpRegion =
result.addRegion();
2465 OpAsmParser::UnresolvedOperand laneId;
2477 result.addAttribute(getWarpSizeAttrName(OperationName(getOperationName(),
2484 llvm::SMLoc inputsOperandsLoc;
2485 SmallVector<OpAsmParser::UnresolvedOperand> inputsOperands;
2486 SmallVector<Type> inputTypes;
2496 if (parser.
resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc,
2507 WarpExecuteOnLane0Op::ensureTerminator(*warpRegion, builder,
result.location);
2515void WarpExecuteOnLane0Op::getSuccessorRegions(
2516 RegionBranchPoint point, SmallVectorImpl<RegionSuccessor> ®ions) {
2518 regions.push_back(RegionSuccessor(getOperation()));
2523 regions.push_back(RegionSuccessor(&getWarpRegion()));
2526ValueRange WarpExecuteOnLane0Op::getSuccessorInputs(RegionSuccessor successor) {
2529void WarpExecuteOnLane0Op::build(OpBuilder &builder, OperationState &
result,
2532 build(builder,
result, resultTypes, laneId, warpSize,
2536void WarpExecuteOnLane0Op::build(OpBuilder &builder, OperationState &
result,
2540 result.addOperands(laneId);
2541 result.addAttribute(getAttributeNames()[0],
2543 result.addTypes(resultTypes);
2544 result.addOperands(args);
2545 assert(args.size() == blockArgTypes.size());
2546 OpBuilder::InsertionGuard guard(builder);
2547 Region *warpRegion =
result.addRegion();
2549 for (
auto [type, arg] : llvm::zip_equal(blockArgTypes, args))
2558 if (expanded == distributed)
2560 auto expandedVecType = llvm::dyn_cast<VectorType>(expanded);
2561 auto distributedVecType = llvm::dyn_cast<VectorType>(distributed);
2562 if (!expandedVecType || !distributedVecType)
2563 return op->
emitOpError(
"expected vector type for distributed operands.");
2564 if (expandedVecType.getRank() != distributedVecType.getRank() ||
2565 expandedVecType.getElementType() != distributedVecType.getElementType())
2567 "expected distributed vectors to have same rank and element type.");
2570 for (
int64_t i = 0, e = expandedVecType.getRank(); i < e; i++) {
2571 int64_t eDim = expandedVecType.getDimSize(i);
2572 int64_t dDim = distributedVecType.getDimSize(i);
2575 if (eDim % dDim != 0)
2577 <<
"expected expanded vector dimension #" << i <<
" (" << eDim
2578 <<
") to be a multipler of the distributed vector dimension ("
2580 scales[i] = eDim / dDim;
2582 if (llvm::product_of(scales) != warpSize)
2584 <<
"incompatible distribution dimensions from " << expandedVecType
2585 <<
" to " << distributedVecType <<
" with warp size = " << warpSize;
2590LogicalResult WarpExecuteOnLane0Op::verify() {
2591 if (getArgs().size() != getWarpRegion().getNumArguments())
2593 "expected same number op arguments and block arguments.");
2594 auto yield = dyn_cast<gpu::YieldOp>(getBody()->getTerminator());
2596 return emitOpError(
"expected body to be terminated with 'gpu.yield'");
2597 if (yield.getNumOperands() != getNumResults())
2599 "expected same number of yield operands and return values.");
2600 int64_t warpSize = getWarpSize();
2601 for (
auto [regionArg, arg] :
2602 llvm::zip_equal(getWarpRegion().getArguments(), getArgs())) {
2604 warpSize, getOperation())))
2607 for (
auto [yieldOperand,
result] :
2608 llvm::zip_equal(yield.getOperands(), getResults())) {
2610 warpSize, getOperation())))
2615bool WarpExecuteOnLane0Op::areTypesCompatible(Type
lhs, Type
rhs) {
2620gpu::YieldOp WarpExecuteOnLane0Op::getTerminator() {
2621 return cast<gpu::YieldOp>(getBody()->getTerminator());
2628void gpu::SubgroupBroadcastOp::inferResultRanges(
2629 ArrayRef<ConstantIntRanges> argRanges,
SetIntRangeFn setResultRange) {
2630 setResultRange(getResult(), argRanges.front());
2634 switch (getBroadcastType()) {
2635 case BroadcastType::first_active_lane:
2639 case BroadcastType::specific_lane:
2643 llvm_unreachable(
"Unknown BroadcastType");
2646LogicalResult gpu::SubgroupBroadcastOp::verify() {
2647 switch (getBroadcastType()) {
2648 case BroadcastType::first_active_lane:
2650 return emitOpError()
2651 <<
"lane can only be specified for `specific_lane` broadcast";
2653 case BroadcastType::specific_lane:
2655 return emitOpError()
2656 <<
"lane must be specified for `specific_lane` broadcast";
2659 llvm_unreachable(
"Unknown BroadcastType");
2662OpFoldResult gpu::SubgroupBroadcastOp::fold(FoldAdaptor ) {
2664 if (
auto prev = getSrc().getDefiningOp<SubgroupBroadcastOp>())
2665 return prev.getResult();
2680KernelMetadataAttr KernelMetadataAttr::get(FunctionOpInterface kernel,
2681 DictionaryAttr metadata) {
2682 assert(kernel &&
"invalid kernel");
2683 return get(kernel.getNameAttr(), kernel.getFunctionType(),
2684 kernel.getAllArgAttrs(), metadata);
2689 FunctionOpInterface kernel,
2690 DictionaryAttr metadata) {
2691 assert(kernel &&
"invalid kernel");
2693 kernel.getAllArgAttrs(), metadata);
2697KernelMetadataAttr::appendMetadata(ArrayRef<NamedAttribute> attrs)
const {
2700 NamedAttrList attrList;
2701 if (DictionaryAttr dict = getMetadata())
2704 return KernelMetadataAttr::get(getName(), getFunctionType(),
getArgAttrs(),
2710 StringAttr name, Type functionType,
2711 ArrayAttr argAttrs, DictionaryAttr metadata) {
2713 return emitError() <<
"the kernel name can't be empty";
2715 if (llvm::any_of(argAttrs, [](Attribute attr) {
2716 return !llvm::isa<DictionaryAttr>(attr);
2719 <<
"all attributes in the array must be a dictionary attribute";
2728KernelTableAttr KernelTableAttr::get(MLIRContext *context,
2729 ArrayRef<KernelMetadataAttr> kernels,
2732 assert((!isSorted || llvm::is_sorted(kernels)) &&
2733 "expected a sorted kernel array");
2735 if (isSorted || llvm::is_sorted(kernels))
2736 return Base::get(context, kernels);
2738 SmallVector<KernelMetadataAttr> kernelsTmp(kernels);
2739 llvm::array_pod_sort(kernelsTmp.begin(), kernelsTmp.end());
2740 return Base::get(context, kernelsTmp);
2743KernelTableAttr KernelTableAttr::getChecked(
2745 ArrayRef<KernelMetadataAttr> kernels,
bool isSorted) {
2747 assert((!isSorted || llvm::is_sorted(kernels)) &&
2748 "expected a sorted kernel array");
2750 if (isSorted || llvm::is_sorted(kernels))
2751 return Base::getChecked(
emitError, context, kernels);
2753 SmallVector<KernelMetadataAttr> kernelsTmp(kernels);
2754 llvm::array_pod_sort(kernelsTmp.begin(), kernelsTmp.end());
2755 return Base::getChecked(
emitError, context, kernelsTmp);
2760 ArrayRef<KernelMetadataAttr> kernels) {
2761 if (kernels.size() < 2)
2764 if (std::adjacent_find(kernels.begin(), kernels.end(),
2765 [](KernelMetadataAttr l, KernelMetadataAttr r) {
2766 return l.getName() == r.getName();
2767 }) != kernels.end()) {
2768 return emitError() <<
"expected all kernels to be uniquely named";
2773KernelMetadataAttr KernelTableAttr::lookup(StringRef key)
const {
2775 return found ? *iterator : KernelMetadataAttr();
2778KernelMetadataAttr KernelTableAttr::lookup(StringAttr key)
const {
2780 return found ? *iterator : KernelMetadataAttr();
2860 return CompilationTarget::Fatbin;
2863std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
2865 std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
options;
2866 llvm::StringSaver stringSaver(
options.first);
2872 if (!opts.empty() && opts.front() ==
'"' && opts.back() ==
'"')
2873 opts.consume_front(
"\""), opts.consume_back(
"\"");
2874 if (!opts.empty() && opts.front() ==
'\'' && opts.back() ==
'\'')
2875 opts.consume_front(
"'"), opts.consume_back(
"'");
2877 llvm::cl::TokenizeWindowsCommandLine(opts, stringSaver,
options.second,
2880 llvm::cl::TokenizeGNUCommandLine(opts, stringSaver,
options.second,
2886std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
2891std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
2893 size_t startPos =
cmdOptions.find(startsWith);
2894 if (startPos == std::string::npos)
2905#include "mlir/Dialect/GPU/IR/GPUOpInterfaces.cpp.inc"
2906#include "mlir/Dialect/GPU/IR/GPUOpsEnums.cpp.inc"
2908#define GET_ATTRDEF_CLASSES
2909#include "mlir/Dialect/GPU/IR/GPUOpsAttributes.cpp.inc"
2911#define GET_OP_CLASSES
2912#include "mlir/Dialect/GPU/IR/GPUOps.cpp.inc"
2914#include "mlir/Dialect/GPU/IR/CompilationAttrInterfaces.cpp.inc"
static void printLaunchFuncOperands(OpAsmPrinter &printer, Operation *, OperandRange operands, TypeRange types)
static ParseResult parseAsyncDependencies(OpAsmParser &parser, Type &asyncTokenType, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &asyncDependencies)
Parses an optional list of async operands with an optional leading keyword.
static void setAttributionAttrs(GPUFuncOp op, unsigned index, DictionaryAttr value, StringAttr attrName)
static void printAttributions(OpAsmPrinter &p, StringRef keyword, ArrayRef< BlockArgument > values, ArrayAttr attributes={})
static LogicalResult verifyDistributedType(Type expanded, Type distributed, int64_t warpSize, Operation *op)
Helper check if the distributed vector type is consistent with the expanded type and distributed size...
static void printAsyncDependencies(OpAsmPrinter &printer, Operation *op, Type asyncTokenType, OperandRange asyncDependencies)
Prints optional async dependencies with its leading keyword.
static ParseResult parseSizeAssignment(OpAsmParser &parser, MutableArrayRef< OpAsmParser::UnresolvedOperand > sizes, MutableArrayRef< OpAsmParser::UnresolvedOperand > regionSizes, MutableArrayRef< OpAsmParser::UnresolvedOperand > indices, StringRef keyword)
static LogicalResult eraseRedundantGpuBarrierOps(BarrierOp op, PatternRewriter &rewriter)
Remove gpu.barrier after gpu.barrier, the threads are already synchronized!
static ParseResult parseOffloadingHandler(OpAsmParser &parser, Attribute &offloadingHandler)
static DictionaryAttr getAttributionAttrs(GPUFuncOp op, unsigned index, StringAttr attrName)
static void printLaunchDimType(OpAsmPrinter &printer, Operation *op, Type dimTy, Value clusterValue, Type clusterXTy, Type clusterYTy, Type clusterZTy)
static bool canMakeGroupOpUniform(Operation *op)
static std::string getSparseHandleKeyword(SparseHandleKind kind)
static LogicalResult verifyKnownLaunchSizeAttr(Operation *op, NamedAttribute attr)
static LogicalResult verifyLaunchAsyncModel(OpTy op)
static ParseResult parseAttributions(OpAsmParser &parser, StringRef keyword, SmallVectorImpl< OpAsmParser::Argument > &args)
Parses a GPU function memory attribution.
static ParseResult parseLaunchDimType(OpAsmParser &parser, Type &dimTy, std::optional< OpAsmParser::UnresolvedOperand > clusterValue, Type &clusterXTy, Type &clusterYTy, Type &clusterZTy)
static void setAttributionAttr(GPUFuncOp op, unsigned index, StringAttr name, Attribute value, StringAttr attrsName)
static ParseResult parseLaunchFuncOperands(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &argNames, SmallVectorImpl< Type > &argTypes)
static void printOffloadingHandler(OpAsmPrinter &printer, Operation *op, Attribute offloadingHandler)
static LogicalResult verifyReduceOpAndType(gpu::AllReduceOperation opName, Type resType)
static void printSizeAssignment(OpAsmPrinter &p, KernelDim3 size, KernelDim3 operands, KernelDim3 ids)
static Attribute getAttributionAttr(GPUFuncOp op, unsigned index, StringAttr name, StringAttr attrsName)
static LogicalResult verifyAttributions(Operation *op, ArrayRef< BlockArgument > attributions, gpu::AddressSpace memorySpace)
Verifies a GPU function memory attribution.
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
static std::string diag(const llvm::Value &value)
static llvm::ManagedStatic< PassManagerOptions > options
template bool mlir::hasSingleEffect< MemoryEffects::Allocate >(Operation *)
static void getDynamicSizes(RankedTensorType tp, ValueRange sizes, SmallVectorImpl< Value > &dynSizes)
Collects the dynamic dimension sizes for tp with the assumption that sizes are the dimension sizes fo...
static sycl::kernel * getKernel(ze_module_handle_t zeModule, const char *name)
#define MLIR_DEFINE_EXPLICIT_TYPE_ID(CLASS_NAME)
ParseResult parseSymbolName(StringAttr &result)
Parse an -identifier and store it (without the '@' symbol) in a string attribute.
@ Paren
Parens surrounding zero or more operands.
@ OptionalSquare
Square brackets supporting zero or more ops, or nothing.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual Location getEncodedSourceLoc(SMLoc loc)=0
Re-encode the given source location as an MLIR location and return it.
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseOptionalColon()=0
Parse a : token if present.
virtual ParseResult parseLSquare()=0
Parse a [ token.
virtual ParseResult parseRSquare()=0
Parse a ] token.
ParseResult parseInteger(IntT &result)
Parse an integer value from the stream.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseDimensionList(SmallVectorImpl< int64_t > &dimensions, bool allowDynamic=true, bool withTrailingX=true)=0
Parse a dimension list of a tensor or memref type.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual ParseResult parseOptionalAttrDictWithKeyword(NamedAttrList &result)=0
Parse a named dictionary into 'result' if the attributes keyword is present.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseColon()=0
Parse a : token.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseOptionalString(std::string *string)=0
Parse a quoted string token if present.
virtual ParseResult parseOptionalLess()=0
Parse a '<' token if present.
virtual ParseResult parseGreater()=0
Parse a '>' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
virtual void printSymbolName(StringRef symbolRef)
Print the given string as a symbol reference, i.e.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
This class is a general helper class for creating context-global objects like types,...
IntegerAttr getI32IntegerAttr(int32_t value)
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
FunctionType getFunctionType(TypeRange inputs, TypeRange results)
IntegerAttr getI64IntegerAttr(int64_t value)
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
StringAttr getStringAttr(const Twine &bytes)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
DictionaryAttr getDictionaryAttr(ArrayRef< NamedAttribute > value)
NamedAttribute getNamedAttr(StringRef name, Attribute val)
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
A symbol reference with a reference path containing a single element.
This class represents a diagnostic that is inflight and set to be reported.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
DictionaryAttr getDictionary(MLIRContext *context) const
Return a dictionary attribute for the underlying dictionary.
void append(StringRef name, Attribute attr)
Add an attribute with the specified name.
NamedAttribute represents a combination of a name and an Attribute value.
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual size_t getNumResults() const =0
Return the number of declared SSA results.
virtual ParseResult parseRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult parseArgumentList(SmallVectorImpl< Argument > &result, Delimiter delimiter=Delimiter::None, bool allowType=false, bool allowAttrs=false)=0
Parse zero or more arguments with a specified surrounding delimiter.
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
static StringRef getOperandSegmentSizeAttr()
This class implements the operand iterators for the Operation class.
Operation is the basic unit of execution within MLIR.
void setInherentAttr(StringAttr name, Attribute value)
Set an inherent attribute by name.
void insertOperands(unsigned index, ValueRange operands)
Insert the given operands into the operand list at the given 'index'.
Block * getBlock()
Returns the operation block that contains this operation.
std::optional< Attribute > getInherentAttr(StringRef name)
Access an inherent attribute by name: returns an empty optional if there is no inherent attribute wit...
Location getLoc()
The source location the operation was defined or derived from.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
AttrClass getDiscardableAttrOfType(StringRef name)
Access a discardable attribute by name and cast it to AttrClass.
MLIRContext * getContext()
Return the context this operation is associated with.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
bool isParent() const
Returns true if branching from the parent op.
bool isOperation() const
Return true if the successor is an operation.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
virtual Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
static Operation * getNearestSymbolTable(Operation *from)
Returns the nearest symbol table from a given operation from.
This class provides an efficient unique identifier for a specific C++ type.
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isSignedInteger() const
Return true if this is a signed integer type (with the specified width).
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
bool isInteger() const
Return true if this is an integer type (with the specified width).
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
user_range getUsers() const
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
static ConcreteType get(MLIRContext *ctx, Args &&...args)
static ConcreteType getChecked(const Location &loc, Args &&...args)
ImplType * getImpl() const
MMAMatrix represents a matrix held by a subgroup for matrix-matrix multiply accumulate operations.
ArrayRef< int64_t > getShape() const
Get shape of the matrix.
static MMAMatrixType get(ArrayRef< int64_t > shape, Type elementType, StringRef operand)
Get MMAMatrixType and verify construction Invariants.
Type getElementType() const
Get elementType of a single element.
static bool isValidElementType(Type elementType)
Check if a type is valid a MMAMatrixType elementType.
static LogicalResult verifyInvariants(function_ref< InFlightDiagnostic()> emitError, ArrayRef< int64_t > shape, Type elementType, StringRef operand)
Verify that shape and elementType are actually allowed for the MMAMatrixType.
StringRef getOperand() const
The general form of operation this type supports is given by the equation C += A*B.
static MMAMatrixType getChecked(function_ref< InFlightDiagnostic()> emitError, ArrayRef< int64_t > shape, Type elementType, StringRef operand)
Get MMAMatrixType at a particular location and verify construction Invariants.
unsigned getNumDims() const
Get number of dims.
This class serves as an opaque interface for passing options to the TargetAttrInterface methods.
function_ref< void(llvm::Module &)> optimizedLlvmIRCallback
Callback invoked with LLVM IR for the device module after LLVM optimizations but before codegen.
function_ref< void(StringRef)> getISACallback() const
Returns the callback invoked with the target ISA for the device, for example PTX assembly.
TypeID getTypeID() const
Returns the typeID.
std::string toolkitPath
Path to the target toolkit.
SymbolTable * getSymbolTable() const
Returns the result of the getSymbolTableCallback callback or a nullptr if no callback was provided.
StringRef getELFSection() const
Returns the ELF section.
StringRef getCmdOptions() const
Returns the command line options.
std::string cmdOptions
An optional set of command line options to be used by the compilation process.
function_ref< void(StringRef)> isaCallback
Callback invoked with the target ISA for the device, for example PTX assembly.
CompilationTarget compilationTarget
Compilation process target format.
std::pair< llvm::BumpPtrAllocator, SmallVector< const char * > > tokenizeCmdOptions() const
Returns a tokenization of the command line options.
function_ref< void(llvm::Module &)> initialLlvmIRCallback
Callback invoked with the initial LLVM IR for the device module.
ArrayRef< Attribute > getLibrariesToLink() const
Returns the LLVM libraries to link to.
TargetOptions(StringRef toolkitPath={}, ArrayRef< Attribute > librariesToLink={}, StringRef cmdOptions={}, StringRef elfSection={}, CompilationTarget compilationTarget=getDefaultCompilationTarget(), function_ref< SymbolTable *()> getSymbolTableCallback={}, function_ref< void(llvm::Module &)> initialLlvmIRCallback={}, function_ref< void(llvm::Module &)> linkedLlvmIRCallback={}, function_ref< void(llvm::Module &)> optimizedLlvmIRCallback={}, function_ref< void(StringRef)> isaCallback={})
Constructor initializing the toolkit path, the list of files to link to, extra command line options,...
function_ref< void(llvm::Module &)> getOptimizedLlvmIRCallback() const
Returns the callback invoked with LLVM IR for the device module after LLVM optimizations but before c...
std::pair< llvm::BumpPtrAllocator, SmallVector< const char * > > tokenizeAndRemoveSuffixCmdOptions(llvm::StringRef startsWith)
Returns a tokenization of the substr of the command line options that starts with startsWith and ends...
StringRef getToolkitPath() const
Returns the toolkit path.
SmallVector< Attribute > librariesToLink
List of files to link with the LLVM module.
function_ref< void(llvm::Module &)> linkedLlvmIRCallback
Callback invoked with LLVM IR for the device module after linking the device libraries.
function_ref< void(llvm::Module &)> getInitialLlvmIRCallback() const
Returns the callback invoked with the initial LLVM IR for the device module.
function_ref< SymbolTable *()> getSymbolTableCallback
Callback for obtaining the parent symbol table of all the GPU modules being serialized.
static CompilationTarget getDefaultCompilationTarget()
Returns the default compilation target: CompilationTarget::Fatbin.
function_ref< void(llvm::Module &)> getLinkedLlvmIRCallback() const
Returns the callback invoked with LLVM IR for the device module after linking the device libraries.
std::string elfSection
ELF Section where the binary needs to be located.
CompilationTarget getCompilationTarget() const
Returns the compilation target.
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto Speculatable
constexpr auto NotSpeculatable
void addArgAndResultAttrs(Builder &builder, OperationState &result, ArrayRef< DictionaryAttr > argAttrs, ArrayRef< DictionaryAttr > resultAttrs, StringAttr argAttrsName, StringAttr resAttrsName)
Adds argument and result attributes, provided as argAttrs and resultAttrs arguments,...
llvm::unique_function< InFlightDiagnostic()> getDefaultDiagnosticEmitFn(MLIRContext *ctx)
Utility method to generate a callback that can be used to generate a diagnostic when checking the con...
ArrayRef< NamedAttribute > getArgAttrs(FunctionOpInterface op, unsigned index)
Return all of the attributes for the argument at 'index'.
ParseResult parseFunctionSignatureWithArguments(OpAsmParser &parser, bool allowVariadic, SmallVectorImpl< OpAsmParser::Argument > &arguments, bool &isVariadic, SmallVectorImpl< Type > &resultTypes, SmallVectorImpl< DictionaryAttr > &resultAttrs)
Parses a function signature using parser.
void printFunctionAttributes(OpAsmPrinter &p, Operation *op, ArrayRef< StringRef > elided={})
Prints the list of function prefixed with the "attributes" keyword.
void printFunctionSignature(OpAsmPrinter &p, FunctionOpInterface op, ArrayRef< Type > argTypes, bool isVariadic, ArrayRef< Type > resultTypes)
Prints the signature of the function-like operation op.
void addAsyncDependency(Operation *op, Value token)
std::pair< IteratorT, bool > findAttrSorted(IteratorT first, IteratorT last, StringRef name)
Using llvm::lower_bound requires an extra string comparison to check whether the returned iterator po...
LogicalResult foldMemRefCast(Operation *op, Value inner=nullptr)
This is a common utility used for patterns of the form "someop(memref.cast) -> someop".
SmallVector< unsigned > getBlockSize(AffineMap dimToLvl)
Given the dimToLvl map, returns the block sizes in a vector.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
llvm::function_ref< void(Value, const ConstantIntRanges &)> SetIntRangeFn
The type of the setResultRanges callback provided to ops implementing InferIntRangeInterface.
LogicalResult verifyDynamicDimensionCount(Operation *op, ShapedType type, ValueRange dynamicSizes)
Verify that the number of dynamic size operands matches the number of dynamic dimensions in the shape...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
auto getChecked(function_ref< InFlightDiagnostic()> emitError, MLIRContext *context, Ts &&...params)
Helper method analogous to get, but uses getChecked when available to allow graceful failure on inval...
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
detail::constant_int_predicate_matcher m_One()
Matches a constant scalar / vector splat / tensor splat integer one.
llvm::TypeSwitch< T, ResultT > TypeSwitch
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
LogicalResult verifyCompatibleShape(ArrayRef< int64_t > shape1, ArrayRef< int64_t > shape2)
Returns success if the given two shapes are compatible.
llvm::function_ref< Fn > function_ref
Simplify the gpu.launch when the range of a thread or block ID is trivially known to be one.
LogicalResult matchAndRewrite(LaunchOp op, PatternRewriter &rewriter) const override
UnresolvedOperand ssaName
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.
Utility class for the GPU dialect to represent triples of Values accessible through ....