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 ) {
616 AllReduceOperationAttr &attr) {
619 std::optional<AllReduceOperation> op =
620 gpu::symbolizeAllReduceOperation(enumStr);
623 attr = AllReduceOperationAttr::get(parser.
getContext(), *op);
629 AllReduceOperationAttr attr) {
638LogicalResult gpu::SubgroupReduceOp::verify() {
640 if (
auto vecTy = dyn_cast<VectorType>(elemType)) {
641 if (vecTy.isScalable())
642 return emitOpError() <<
"is not compatible with scalable vector types";
644 elemType = vecTy.getElementType();
647 gpu::AllReduceOperation opName = getOp();
649 return emitError() <<
'`' << gpu::stringifyAllReduceOperation(opName)
650 <<
"` reduction operation is not compatible with type "
654 auto clusterSize = getClusterSize();
656 uint32_t size = *clusterSize;
657 if (!llvm::isPowerOf2_32(size)) {
659 <<
" is not a power of two";
663 uint32_t stride = getClusterStride();
664 if (stride != 1 && !clusterSize) {
665 return emitOpError() <<
"cluster stride can only be specified if cluster "
668 if (!llvm::isPowerOf2_32(stride)) {
669 return emitOpError() <<
"cluster stride " << stride
670 <<
" is not a power of two";
676OpFoldResult gpu::SubgroupReduceOp::fold(FoldAdaptor ) {
677 if (getClusterSize() == 1)
694 if (!op->template hasTrait<OpTrait::AttrSizedOperandSegments>())
698 auto sizeAttr = op->template getAttrOfType<DenseI32ArrayAttr>(attrName);
716 Value getBlockSizeZ,
Value dynamicSharedMemorySize,
724 if (!workgroupAttributions.empty())
726 getWorkgroupAttributionsAttrName(
result.name),
730 result.addOperands(asyncDependencies);
735 result.addOperands({gridSizeX, gridSizeY, gridSizeZ, getBlockSizeX,
736 getBlockSizeY, getBlockSizeZ});
738 result.addOperands(clusterSizeX);
740 result.addOperands(clusterSizeY);
742 result.addOperands(clusterSizeZ);
743 if (dynamicSharedMemorySize)
744 result.addOperands(dynamicSharedMemorySize);
746 result.addOperands(asyncObject);
750 result.addAttribute(getModuleAttrName(
result.name), module);
752 result.addAttribute(getFunctionAttrName(
result.name), function);
760 for (
unsigned i = 0; i < kNumConfigRegionAttributes; ++i)
763 for (
Type argTy : workgroupAttributions)
765 for (
Type argTy : privateAttributions)
769 segmentSizes.front() = asyncDependencies.size();
770 segmentSizes[7] = clusterSizeX ? 1 : 0;
771 segmentSizes[8] = clusterSizeY ? 1 : 0;
772 segmentSizes[9] = clusterSizeZ ? 1 : 0;
773 segmentSizes[10] = dynamicSharedMemorySize ? 1 : 0;
774 segmentSizes[11] = asyncObject ? 1 : 0;
775 result.addAttribute(getOperandSegmentSizeAttr(),
780 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
781 auto args = getBody().getArguments();
786 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
787 auto args = getBody().getArguments();
792 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
793 auto args = getBody().getArguments();
798 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
799 auto args = getBody().getArguments();
800 return KernelDim3{args[9], args[10], args[11]};
803std::optional<KernelDim3> LaunchOp::getClusterIds() {
804 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
805 if (!hasClusterSize())
807 auto args = getBody().getArguments();
808 return KernelDim3{args[12], args[13], args[14]};
811std::optional<KernelDim3> LaunchOp::getClusterSize() {
812 assert(!getBody().empty() &&
"LaunchOp body must not be empty.");
813 if (!hasClusterSize())
815 auto args = getBody().getArguments();
816 return KernelDim3{args[15], args[16], args[17]};
819KernelDim3 LaunchOp::getGridSizeOperandValues() {
820 auto operands = getOperands().drop_front(getAsyncDependencies().size());
821 return KernelDim3{operands[0], operands[1], operands[2]};
824KernelDim3 LaunchOp::getBlockSizeOperandValues() {
825 auto operands = getOperands().drop_front(getAsyncDependencies().size());
826 return KernelDim3{operands[3], operands[4], operands[5]};
829std::optional<KernelDim3> LaunchOp::getClusterSizeOperandValues() {
830 auto operands = getOperands().drop_front(getAsyncDependencies().size());
831 if (!hasClusterSize())
833 return KernelDim3{operands[6], operands[7], operands[8]};
836template <
typename OpTy>
838 if (!op.getAsyncDependencies().empty() && !op.getAsyncToken())
839 return op.emitOpError(
"dependency operands require the dependency-based "
840 "async model i.e. returning a token");
841 if (op.getAsyncToken() && op.getAsyncObject())
842 return op.emitOpError(
"stream-based and dependency-based async models are "
843 "mutually exclusive");
844 if (op.getNumResults() == 0 && op.getAsyncToken())
845 return op.emitOpError(
"needs to be named when async keyword is specified");
849LogicalResult LaunchOp::verify() {
853 if (!(hasClusterSize()) &&
854 (getClusterSizeX() || getClusterSizeY() || getClusterSizeZ()))
855 return emitOpError() <<
"cluster size must be all present";
859LogicalResult LaunchOp::verifyRegions() {
863 if (getBody().empty()) {
866 unsigned actualNumRegionArgs = getBody().getNumArguments();
867 unsigned expectedNumRegionArgs =
868 getNumConfigRegionAttributes() + getNumWorkgroupAttributions();
869 if (actualNumRegionArgs < expectedNumRegionArgs) {
871 << expectedNumRegionArgs <<
" region arguments, but got "
872 << actualNumRegionArgs;
877 GPUDialect::getWorkgroupAddressSpace())) ||
879 GPUDialect::getPrivateAddressSpace())))
884 for (
Block &block : getBody()) {
887 if (block.back().getNumSuccessors() != 0)
889 if (!isa<gpu::TerminatorOp>(&block.back())) {
892 .append(
"expected '", gpu::TerminatorOp::getOperationName(),
893 "' or a terminator with successors")
894 .attachNote(getLoc())
895 .append(
"in '", LaunchOp::getOperationName(),
"' body region");
908 p <<
'(' << ids.
x <<
", " << ids.
y <<
", " << ids.
z <<
") in (";
909 p << size.
x <<
" = " << operands.
x <<
", ";
910 p << size.
y <<
" = " << operands.
y <<
", ";
911 p << size.
z <<
" = " << operands.
z <<
')';
914void LaunchOp::print(OpAsmPrinter &p) {
915 if (
auto asyncObject = getAsyncObject()) {
916 p <<
" <" << asyncObject <<
" : " << asyncObject.
getType() <<
">";
918 if (getAsyncToken()) {
920 if (!getAsyncDependencies().empty())
921 p <<
" [" << getAsyncDependencies() <<
']';
924 if (hasClusterSize()) {
925 p <<
' ' << getClustersKeyword();
927 getClusterSizeOperandValues().value(),
928 getClusterIds().value());
930 p <<
' ' << getBlocksKeyword();
933 p <<
' ' << getThreadsKeyword();
936 if (getDynamicSharedMemorySize())
937 p <<
' ' << getDynamicSharedMemorySizeKeyword() <<
' '
938 << getDynamicSharedMemorySize();
941 StringRef moduleAttrName = getModuleAttrName();
942 if (
auto module = getModule()) {
943 p <<
' ' << moduleAttrName <<
'(';
948 StringRef functionAttrName = getFunctionAttrName();
949 if (
auto function = getFunction()) {
950 p <<
' ' << functionAttrName <<
'(';
955 if (getCooperative())
965 LaunchOp::getOperandSegmentSizeAttr(),
966 getWorkgroupAttributionsAttrName(),
967 getCooperativeAttrName(), moduleAttrName,
983 assert(
indices.size() == 3 &&
"space for three indices expected");
990 if (args.size() != 3) {
992 << keyword <<
" expects 3 arguments, but got " << args.size();
994 std::move(args.begin(), args.end(),
indices.begin());
996 for (
int i = 0; i < 3; ++i) {
1019ParseResult LaunchOp::parse(OpAsmParser &parser, OperationState &
result) {
1021 SmallVector<OpAsmParser::UnresolvedOperand, LaunchOp::kNumConfigOperands>
1022 sizes(LaunchOp::kNumConfigOperands);
1025 SmallVector<OpAsmParser::UnresolvedOperand, 16> regionArgs(
1026 LaunchOp::kNumConfigRegionAttributes);
1029 OpAsmParser::UnresolvedOperand asyncObjectOperand;
1030 Type asyncObjectType;
1031 bool hasAsyncObject =
false;
1033 hasAsyncObject =
true;
1040 SmallVector<OpAsmParser::UnresolvedOperand, 4> asyncDependencies;
1041 Type asyncTokenType;
1048 if (!asyncTokenType)
1051 "gpu.launch requires 'async' keyword to return a value");
1052 result.types.push_back(asyncTokenType);
1055 bool hasCluster =
false;
1059 regionArgs.resize(18);
1061 MutableArrayRef<OpAsmParser::UnresolvedOperand> sizesRef(sizes);
1062 MutableArrayRef<OpAsmParser::UnresolvedOperand> regionArgsRef(regionArgs);
1068 parser, sizesRef.drop_front(6), regionArgsRef.slice(15, 3),
1069 regionArgsRef.slice(12, 3), LaunchOp::getClustersKeyword()))
1077 if (parser.
parseKeyword(LaunchOp::getBlocksKeyword()) ||
1079 regionArgsRef.slice(6, 3), regionArgsRef.slice(0, 3),
1080 LaunchOp::getBlocksKeyword()) ||
1083 regionArgsRef.slice(9, 3), regionArgsRef.slice(3, 3),
1084 LaunchOp::getThreadsKeyword()) ||
1089 OpAsmParser::UnresolvedOperand dynamicSharedMemorySize;
1090 bool hasDynamicSharedMemorySize =
false;
1092 LaunchOp::getDynamicSharedMemorySizeKeyword())) {
1093 hasDynamicSharedMemorySize =
true;
1103 asyncObjectType,
result.operands))
1107 StringRef moduleAttrName = getModuleAttrName(
result.name);
1109 FlatSymbolRefAttr moduleSymbol;
1117 StringRef functionAttrName = getFunctionAttrName(
result.name);
1119 FlatSymbolRefAttr funcSymbol;
1135 SmallVector<Type, LaunchOp::kNumConfigRegionAttributes> dataTypes(
1136 LaunchOp::kNumConfigRegionAttributes + 6, index);
1138 SmallVector<OpAsmParser::Argument> regionArguments;
1139 for (
auto ssaValueAndType : llvm::zip(regionArgs, dataTypes)) {
1140 OpAsmParser::Argument arg;
1141 arg.
ssaName = std::get<0>(ssaValueAndType);
1142 arg.
type = std::get<1>(ssaValueAndType);
1143 regionArguments.push_back(arg);
1154 unsigned numWorkgroupAttrs = regionArguments.size() -
1155 LaunchOp::kNumConfigRegionAttributes -
1156 (hasCluster ? 6 : 0);
1157 if (numWorkgroupAttrs != 0)
1158 result.addAttribute(LaunchOp::getWorkgroupAttributionsAttrName(
result.name),
1169 Region *body =
result.addRegion();
1174 SmallVector<int32_t, 12> segmentSizes(12, 1);
1175 segmentSizes.front() = asyncDependencies.size();
1178 segmentSizes[7] = 0;
1179 segmentSizes[8] = 0;
1180 segmentSizes[9] = 0;
1182 segmentSizes[10] = hasDynamicSharedMemorySize ? 1 : 0;
1183 segmentSizes[11] = hasAsyncObject ? 1 : 0;
1184 result.addAttribute(LaunchOp::getOperandSegmentSizeAttr(),
1198 bool simplified =
false;
1199 auto constPropIdUses = [&](
Value id,
Value size) {
1203 if (
id.getUses().empty())
1215 constPropIdUses(op.getBlockIds().x, op.getGridSizeX());
1216 constPropIdUses(op.getBlockIds().y, op.getGridSizeY());
1217 constPropIdUses(op.getBlockIds().z, op.getGridSizeZ());
1218 constPropIdUses(op.getThreadIds().x, op.getBlockSizeX());
1219 constPropIdUses(op.getThreadIds().y, op.getBlockSizeY());
1220 constPropIdUses(op.getThreadIds().z, op.getBlockSizeZ());
1226void LaunchOp::getCanonicalizationPatterns(RewritePatternSet &rewrites,
1227 MLIRContext *context) {
1228 rewrites.
add<FoldLaunchArguments>(context);
1233BlockArgument LaunchOp::addWorkgroupAttribution(Type type, Location loc) {
1234 int64_t cur = getWorkgroupAttributions().value_or(0);
1235 setWorkgroupAttributions(std::optional<int64_t>(cur + 1));
1236 return getBody().insertArgument(
1237 getNumConfigRegionAttributes() +
static_cast<unsigned>(cur), type, loc);
1242BlockArgument LaunchOp::addPrivateAttribution(Type type, Location loc) {
1245 return getBody().addArgument(type, loc);
1252void LaunchFuncOp::build(OpBuilder &builder, OperationState &
result,
1253 SymbolRefAttr kernelSymbol,
KernelDim3 gridSize,
1254 KernelDim3 getBlockSize, Value dynamicSharedMemorySize,
1255 ValueRange kernelOperands, Type asyncTokenType,
1256 ValueRange asyncDependencies, Value asyncObject,
1257 std::optional<KernelDim3> clusterSize) {
1258 assert(kernelSymbol.getNestedReferences().size() == 1 &&
1259 "expected a symbol reference with a single nested reference");
1260 result.addOperands(asyncDependencies);
1267 if (clusterSize.has_value())
1268 result.addOperands({clusterSize->x, clusterSize->y, clusterSize->z});
1269 if (dynamicSharedMemorySize)
1270 result.addOperands(dynamicSharedMemorySize);
1271 result.addOperands(kernelOperands);
1273 result.addOperands(asyncObject);
1275 Properties &prop =
result.getOrAddProperties<Properties>();
1276 prop.kernel = kernelSymbol;
1277 size_t segmentSizesLen = std::size(prop.operandSegmentSizes);
1279 llvm::fill(prop.operandSegmentSizes, 1);
1280 prop.operandSegmentSizes[0] = asyncDependencies.size();
1281 if (!clusterSize.has_value()) {
1282 prop.operandSegmentSizes[segmentSizesLen - 4] = 0;
1283 prop.operandSegmentSizes[segmentSizesLen - 5] = 0;
1284 prop.operandSegmentSizes[segmentSizesLen - 6] = 0;
1286 prop.operandSegmentSizes[segmentSizesLen - 3] =
1287 dynamicSharedMemorySize ? 1 : 0;
1288 prop.operandSegmentSizes[segmentSizesLen - 2] =
1289 static_cast<int32_t
>(kernelOperands.size());
1290 prop.operandSegmentSizes[segmentSizesLen - 1] = asyncObject ? 1 : 0;
1293void LaunchFuncOp::build(OpBuilder &builder, OperationState &
result,
1295 KernelDim3 getBlockSize, Value dynamicSharedMemorySize,
1296 ValueRange kernelOperands, Type asyncTokenType,
1297 ValueRange asyncDependencies, Value asyncObject,
1298 std::optional<KernelDim3> clusterSize) {
1299 auto kernelModule = kernelFunc->getParentOfType<GPUModuleOp>();
1301 SymbolRefAttr::get(kernelModule.getNameAttr(),
1302 {SymbolRefAttr::get(kernelFunc.getNameAttr())});
1303 build(builder,
result, kernelSymbol, gridSize, getBlockSize,
1304 dynamicSharedMemorySize, kernelOperands, asyncTokenType,
1305 asyncDependencies, asyncObject, clusterSize);
1308StringAttr LaunchFuncOp::getKernelModuleName() {
1312StringAttr LaunchFuncOp::getKernelName() {
1316unsigned LaunchFuncOp::getNumKernelOperands() {
1317 return getKernelOperands().size();
1320Value LaunchFuncOp::getKernelOperand(
unsigned i) {
1321 return getKernelOperands()[i];
1324KernelDim3 LaunchFuncOp::getGridSizeOperandValues() {
1325 auto operands = getOperands().drop_front(getAsyncDependencies().size());
1326 return KernelDim3{operands[0], operands[1], operands[2]};
1329KernelDim3 LaunchFuncOp::getBlockSizeOperandValues() {
1330 auto operands = getOperands().drop_front(getAsyncDependencies().size());
1331 return KernelDim3{operands[3], operands[4], operands[5]};
1334KernelDim3 LaunchFuncOp::getClusterSizeOperandValues() {
1335 assert(hasClusterSize() &&
1336 "cluster size is not set, check hasClusterSize() first");
1337 auto operands = getOperands().drop_front(getAsyncDependencies().size());
1338 return KernelDim3{operands[6], operands[7], operands[8]};
1341LogicalResult LaunchFuncOp::verify() {
1345 auto module = (*this)->getParentOfType<ModuleOp>();
1347 return emitOpError(
"expected to belong to a module");
1349 if (!module->getAttrOfType<UnitAttr>(
1350 GPUDialect::getContainerModuleAttrName()))
1351 return emitOpError(
"expected the closest surrounding module to have the '" +
1352 GPUDialect::getContainerModuleAttrName() +
1355 if (hasClusterSize()) {
1356 if (getClusterSizeY().
getType() != getClusterSizeX().
getType() ||
1359 <<
"expects types of the cluster dimensions must be the same";
1366LaunchFuncOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
1367 LaunchFuncOp launchOp = *
this;
1370 if (isa<GPUModuleOp>(table))
1375 if (!launchOp->getParentOp() ||
1376 launchOp->getParentOp()->getParentOp() != table)
1381 if (!launchOp->getAttrOfType<SymbolRefAttr>(
1382 LaunchFuncOp::getKernelAttrName(launchOp->getName())))
1386 StringAttr kernelContainerName = launchOp.getKernelModuleName();
1387 Operation *kernelContainer =
1389 if (!kernelContainer)
1391 <<
"kernel container '" << kernelContainerName.getValue()
1392 <<
"' is undefined";
1395 if (isa<BinaryOp>(kernelContainer))
1398 auto kernelModule = dyn_cast<GPUModuleOp>(kernelContainer);
1400 return launchOp.emitOpError()
1401 <<
"kernel module '" << kernelContainerName.getValue()
1402 <<
"' is undefined";
1406 kernelModule, launchOp.getKernelName());
1409 << launchOp.getKernel() <<
"' is undefined";
1410 auto kernelConvertedFunction = dyn_cast<FunctionOpInterface>(kernelFunc);
1411 if (!kernelConvertedFunction) {
1412 InFlightDiagnostic
diag = launchOp.emitOpError()
1413 <<
"referenced kernel '" << launchOp.getKernel()
1414 <<
"' is not a function";
1415 diag.attachNote(kernelFunc->
getLoc()) <<
"see the kernel definition here";
1419 if (!GPUDialect::isKernel(kernelFunc))
1420 return launchOp.emitOpError(
"kernel function is missing the '")
1421 << GPUDialect::getKernelFuncAttrName() <<
"' attribute";
1426 auto kernelGPUFunction = dyn_cast<gpu::GPUFuncOp>(kernelFunc);
1427 if (!kernelGPUFunction)
1430 unsigned actualNumArguments = launchOp.getNumKernelOperands();
1431 unsigned expectedNumArguments = kernelGPUFunction.getNumArguments();
1432 if (expectedNumArguments != actualNumArguments)
1433 return launchOp.emitOpError(
"got ")
1434 << actualNumArguments <<
" kernel operands but expected "
1435 << expectedNumArguments;
1437 FunctionType functionType = kernelGPUFunction.getFunctionType();
1438 for (
unsigned i = 0; i < expectedNumArguments; ++i) {
1439 if (launchOp.getKernelOperand(i).getType() != functionType.getInput(i)) {
1440 return launchOp.emitOpError(
"type of function argument ")
1441 << i <<
" does not match";
1450 std::optional<OpAsmParser::UnresolvedOperand> clusterValue,
1451 Type &clusterXTy,
Type &clusterYTy,
Type &clusterZTy) {
1458 if (clusterValue.has_value()) {
1459 clusterXTy = clusterYTy = clusterZTy = dimTy;
1466 Type clusterYTy,
Type clusterZTy) {
1468 printer <<
": " << dimTy;
1478 auto parseElement = [&]() -> ParseResult {
1479 return failure(parser.
parseOperand(argNames.emplace_back()) ||
1484 parseElement,
" in argument list");
1489 if (operands.empty())
1492 llvm::interleaveComma(llvm::zip_equal(operands, types), printer,
1493 [&](
const auto &pair) {
1494 auto [operand, type] = pair;
1495 printer << operand <<
" : " << type;
1504void ShuffleOp::build(OpBuilder &builder, OperationState &
result, Value value,
1505 int32_t offset, int32_t width, ShuffleMode mode) {
1506 build(builder,
result, value,
1507 arith::ConstantOp::create(builder,
result.location,
1509 arith::ConstantOp::create(builder,
result.location,
1518LogicalResult RotateOp::verify() {
1519 uint32_t offset = getOffset();
1520 uint32_t width = getWidth();
1522 if (offset >= width) {
1523 return emitOpError() <<
"offset must be in the range [0, " << width <<
")";
1533LogicalResult BarrierOp::verify() {
1534 BarrierScope scope = getScope();
1536 if (getNamedBarrier() && scope != BarrierScope::Workgroup)
1537 return emitOpError(
"named barriers require workgroup scope");
1545 auto nextOp = dyn_cast_or_null<BarrierOp>(op->getNextNode());
1550 if (op.getScope() != nextOp.getScope())
1554 if (op.getNamedBarrier() != nextOp.getNamedBarrier())
1557 std::optional<ArrayAttr> thisMemfence = op.getAddressSpaces();
1558 std::optional<ArrayAttr> nextMemfence = nextOp.getAddressSpaces();
1562 if (!nextMemfence) {
1563 op.removeAddressSpacesAttr();
1567 if (*thisMemfence == *nextMemfence) {
1571 llvm::SmallSetVector<Attribute, 4> mergedSpaces;
1573 mergedSpaces.insert(attr);
1575 mergedSpaces.insert(attr);
1576 op.setAddressSpacesAttr(rewriter.
getArrayAttr(mergedSpaces.takeVector()));
1584void BarrierOp::getCanonicalizationPatterns(RewritePatternSet &results,
1585 MLIRContext *context) {
1589void BarrierOp::build(mlir::OpBuilder &odsBuilder,
1590 mlir::OperationState &odsState,
1591 std::optional<AddressSpace> addressSpace) {
1595 AddressSpaceAttr::get(odsBuilder.
getContext(), addressSpace.value()));
1597 odsBuilder, odsState, addressSpacesAttr, Value{},
1598 BarrierScopeAttr::get(odsBuilder.
getContext(), BarrierScope::Workgroup));
1605void BarrierOp::build(OpBuilder &builder, OperationState &odsState,
1606 Value memrefToFence) {
1607 std::optional<AddressSpace> addrSpaceToFence;
1608 if (
auto memrefType = dyn_cast<BaseMemRefType>(memrefToFence.
getType()))
1609 if (
auto addrSpaceAttr = dyn_cast_if_present<gpu::AddressSpaceAttr>(
1610 memrefType.getMemorySpace()))
1611 addrSpaceToFence = addrSpaceAttr.getValue();
1612 return build(builder, odsState, addrSpaceToFence);
1621BlockArgument GPUFuncOp::addWorkgroupAttribution(Type type, Location loc) {
1622 int64_t cur = getWorkgroupAttributions().value_or(0);
1623 setWorkgroupAttributions(std::optional<int64_t>(cur + 1));
1624 return getBody().insertArgument(
1625 getFunctionType().getNumInputs() +
static_cast<unsigned>(cur), type, loc);
1630BlockArgument GPUFuncOp::addPrivateAttribution(Type type, Location loc) {
1633 return getBody().addArgument(type, loc);
1636void GPUFuncOp::build(OpBuilder &builder, OperationState &
result,
1637 StringRef name, FunctionType type,
1640 ArrayRef<NamedAttribute> attrs) {
1641 OpBuilder::InsertionGuard g(builder);
1645 result.addAttribute(getFunctionTypeAttrName(
result.name),
1646 TypeAttr::get(type));
1647 result.addAttribute(getWorkgroupAttributionsAttrName(
result.name),
1649 result.addAttributes(attrs);
1650 Region *body =
result.addRegion();
1654 for (Type argTy : type.getInputs())
1656 for (Type argTy : workgroupAttributions)
1658 for (Type argTy : privateAttributions)
1677 size_t existingArgs = args.size();
1684 bool hadAttrs = llvm::any_of(
ArrayRef(args).drop_front(existingArgs),
1689 attributionAttrs =
nullptr;
1695 for (
const auto &argument :
ArrayRef(args).drop_front(existingArgs)) {
1696 if (!argument.attrs)
1699 attributionAttrsVec.push_back(argument.attrs);
1701 attributionAttrs = builder.
getArrayAttr(attributionAttrsVec);
1710ParseResult GPUFuncOp::parse(OpAsmParser &parser, OperationState &
result) {
1711 SmallVector<OpAsmParser::Argument> entryArgs;
1712 SmallVector<DictionaryAttr> resultAttrs;
1713 SmallVector<Type> resultTypes;
1717 StringAttr nameAttr;
1724 parser,
false, entryArgs, isVariadic, resultTypes,
1728 if (!entryArgs.empty() && entryArgs[0].ssaName.name.empty())
1729 return parser.
emitError(signatureLocation)
1730 <<
"gpu.func requires named arguments";
1736 SmallVector<Type> argTypes;
1737 for (
auto &arg : entryArgs)
1738 argTypes.push_back(arg.
type);
1740 result.addAttribute(getFunctionTypeAttrName(
result.name),
1741 TypeAttr::get(type));
1744 builder,
result, entryArgs, resultAttrs, getArgAttrsAttrName(
result.name),
1745 getResAttrsAttrName(
result.name));
1747 Attribute workgroupAttributionAttrs;
1750 entryArgs, workgroupAttributionAttrs)))
1755 unsigned numWorkgroupAttrs = entryArgs.size() - type.getNumInputs();
1756 if (numWorkgroupAttrs != 0)
1758 GPUFuncOp::getWorkgroupAttributionsAttrName(
result.name),
1760 if (workgroupAttributionAttrs)
1761 result.addAttribute(GPUFuncOp::getWorkgroupAttribAttrsAttrName(
result.name),
1762 workgroupAttributionAttrs);
1764 Attribute privateAttributionAttrs;
1767 entryArgs, privateAttributionAttrs)))
1769 if (privateAttributionAttrs)
1770 result.addAttribute(GPUFuncOp::getPrivateAttribAttrsAttrName(
result.name),
1771 privateAttributionAttrs);
1775 result.addAttribute(GPUFuncOp::getKernelAttrName(
result.name),
1784 auto *body =
result.addRegion();
1788void GPUFuncOp::print(OpAsmPrinter &p) {
1792 FunctionType type = getFunctionType();
1798 getWorkgroupAttribAttrs().value_or(
nullptr));
1800 getPrivateAttribAttrs().value_or(
nullptr));
1802 p <<
' ' << getKernelKeyword();
1806 {getWorkgroupAttributionsAttrName(), getKernelAttrName(),
1807 GPUDialect::getKernelFuncAttrName(), getFunctionTypeAttrName(),
1808 getArgAttrsAttrName(), getResAttrsAttrName(),
1809 getWorkgroupAttribAttrsAttrName(), getPrivateAttribAttrsAttrName()});
1815 StringAttr attrName) {
1816 auto allAttrs = llvm::dyn_cast_or_null<ArrayAttr>(op->getAttr(attrName));
1817 if (!allAttrs ||
index >= allAttrs.size())
1818 return DictionaryAttr();
1819 return llvm::cast<DictionaryAttr>(allAttrs[
index]);
1822DictionaryAttr GPUFuncOp::getworkgroupAttributionAttrs(
unsigned index) {
1826DictionaryAttr GPUFuncOp::getPrivateAttributionAttrs(
unsigned index) {
1831 DictionaryAttr value, StringAttr attrName) {
1833 auto allAttrs = llvm::dyn_cast_or_null<ArrayAttr>(op->getAttr(attrName));
1836 elements.append(allAttrs.begin(), allAttrs.end());
1837 while (elements.size() <=
index)
1838 elements.push_back(DictionaryAttr::get(ctx));
1840 elements[
index] = DictionaryAttr::get(ctx);
1842 elements[
index] = value;
1843 ArrayAttr newValue = ArrayAttr::get(ctx, elements);
1844 op->setAttr(attrName, newValue);
1847void GPUFuncOp::setworkgroupAttributionAttrs(
unsigned index,
1848 DictionaryAttr value) {
1852void GPUFuncOp::setPrivateAttributionAttrs(
unsigned int index,
1853 DictionaryAttr value) {
1858 StringAttr name, StringAttr attrsName) {
1862 return dict.get(name);
1865Attribute GPUFuncOp::getWorkgroupAttributionAttr(
unsigned index,
1867 assert(index < getNumWorkgroupAttributions() &&
1868 "index must map to a workgroup attribution");
1870 getWorkgroupAttribAttrsAttrName());
1873Attribute GPUFuncOp::getPrivateAttributionAttr(
unsigned index,
1875 assert(index < getNumPrivateAttributions() &&
1876 "index must map to a private attribution");
1878 getPrivateAttribAttrsAttrName());
1882 Attribute value, StringAttr attrsName) {
1887 elems.append(oldDict.getValue().begin(), oldDict.getValue().end());
1890 bool mustSort =
true;
1891 for (
unsigned i = 0, e = elems.size(); i < e; ++i) {
1892 if (elems[i].getName() == name) {
1895 std::swap(elems[i], elems[elems.size() - 1]);
1907 elems.emplace_back(name, value);
1910 DictionaryAttr::sortInPlace(elems);
1912 auto newDict = DictionaryAttr::getWithSorted(ctx, elems);
1916void GPUFuncOp::setWorkgroupAttributionAttr(
unsigned index, StringAttr name,
1918 assert(index < getNumWorkgroupAttributions() &&
1919 "index must map to a workgroup attribution");
1921 getWorkgroupAttribAttrsAttrName());
1924void GPUFuncOp::setPrivateAttributionAttr(
unsigned index, StringAttr name,
1926 assert(index < getNumPrivateAttributions() &&
1927 "index must map to a private attribution");
1929 getPrivateAttribAttrsAttrName());
1932LogicalResult GPUFuncOp::verifyType() {
1933 if (isKernel() && getFunctionType().getNumResults() != 0)
1934 return emitOpError() <<
"expected void return type for kernel function";
1940LogicalResult GPUFuncOp::verifyBody() {
1942 return emitOpError() <<
"expected body with at least one block";
1943 unsigned numFuncArguments = getNumArguments();
1944 unsigned numWorkgroupAttributions = getNumWorkgroupAttributions();
1945 unsigned numBlockArguments = front().getNumArguments();
1946 if (numBlockArguments < numFuncArguments + numWorkgroupAttributions)
1948 << numFuncArguments + numWorkgroupAttributions
1949 <<
" arguments to body region";
1951 ArrayRef<Type> funcArgTypes = getFunctionType().getInputs();
1952 for (
unsigned i = 0; i < numFuncArguments; ++i) {
1953 Type blockArgType = front().getArgument(i).getType();
1954 if (funcArgTypes[i] != blockArgType)
1955 return emitOpError() <<
"expected body region argument #" << i
1956 <<
" to be of type " << funcArgTypes[i] <<
", got "
1961 GPUDialect::getWorkgroupAddressSpace())) ||
1963 GPUDialect::getPrivateAddressSpace())))
1973LogicalResult gpu::ReturnOp::verify() {
1974 GPUFuncOp function = (*this)->getParentOfType<GPUFuncOp>();
1976 FunctionType funType = function.getFunctionType();
1978 if (funType.getNumResults() != getOperands().size())
1980 .append(
"expected ", funType.getNumResults(),
" result operands")
1981 .attachNote(function.getLoc())
1982 .append(
"return type declared here");
1984 for (
const auto &pair : llvm::enumerate(
1985 llvm::zip(function.getFunctionType().getResults(), getOperands()))) {
1986 auto [type, operand] = pair.value();
1987 if (type != operand.getType())
1988 return emitOpError() <<
"unexpected type `" << operand.getType()
1989 <<
"' for operand #" << pair.index();
1998void GPUModuleOp::build(OpBuilder &builder, OperationState &
result,
2000 Attribute offloadingHandler) {
2001 result.addRegion()->emplaceBlock();
2002 Properties &props =
result.getOrAddProperties<Properties>();
2004 props.targets = targets;
2006 props.offloadingHandler = offloadingHandler;
2009void GPUModuleOp::build(OpBuilder &builder, OperationState &
result,
2010 StringRef name, ArrayRef<Attribute> targets,
2011 Attribute offloadingHandler) {
2012 build(builder,
result, name,
2017bool GPUModuleOp::hasTarget(Attribute
target) {
2018 if (
ArrayAttr targets = getTargetsAttr())
2019 return llvm::count(targets.getValue(),
target);
2023void GPUModuleOp::setTargets(ArrayRef<TargetAttrInterface> targets) {
2024 ArrayAttr &targetsAttr = getProperties().targets;
2025 SmallVector<Attribute> targetsVector(targets);
2026 targetsAttr = ArrayAttr::get(
getContext(), targetsVector);
2029LogicalResult GPUModuleOp::verify() {
2030 auto targets = getOperation()->getAttrOfType<
ArrayAttr>(
"targets");
2035 for (
auto target : targets) {
2036 if (
auto verifyTargetAttr =
2037 llvm::dyn_cast<TargetAttrVerifyInterface>(
target)) {
2038 if (verifyTargetAttr.verifyTarget(getOperation()).
failed())
2048void BinaryOp::build(OpBuilder &builder, OperationState &
result, StringRef name,
2049 Attribute offloadingHandler,
ArrayAttr objects) {
2050 auto &properties =
result.getOrAddProperties<Properties>();
2053 properties.objects = objects;
2054 if (offloadingHandler)
2055 properties.offloadingHandler = offloadingHandler;
2057 properties.offloadingHandler = builder.
getAttr<SelectObjectAttr>(
nullptr);
2060void BinaryOp::build(OpBuilder &builder, OperationState &
result, StringRef name,
2061 Attribute offloadingHandler, ArrayRef<Attribute> objects) {
2062 build(builder,
result, name, offloadingHandler,
2074 if (!offloadingHandler)
2081 if (offloadingHandler != SelectObjectAttr::get(op->
getContext(),
nullptr))
2082 printer <<
'<' << offloadingHandler <<
'>';
2089LogicalResult MemcpyOp::verify() {
2090 auto srcType = getSrc().getType();
2091 auto dstType = getDst().getType();
2094 return emitOpError(
"arguments have incompatible element type");
2097 return emitOpError(
"arguments have incompatible shape");
2106struct EraseTrivialCopyOp :
public OpRewritePattern<MemcpyOp> {
2107 using OpRewritePattern<MemcpyOp>::OpRewritePattern;
2109 LogicalResult matchAndRewrite(MemcpyOp op,
2110 PatternRewriter &rewriter)
const override {
2111 Value dest = op.getDst();
2120 if (llvm::any_of(dest.
getUsers(), [op, dest](Operation *user) {
2121 return user != op &&
2122 !hasSingleEffect<MemoryEffects::Free>(user, dest);
2128 if (op.getAsyncDependencies().size() > 1 ||
2129 ((op.getAsyncDependencies().empty() && op.getAsyncToken()) ||
2130 (!op.getAsyncDependencies().empty() && !op.getAsyncToken())))
2132 rewriter.
replaceOp(op, op.getAsyncDependencies());
2139void MemcpyOp::getCanonicalizationPatterns(RewritePatternSet &results,
2140 MLIRContext *context) {
2141 results.
add<EraseTrivialCopyOp>(context);
2148LogicalResult SubgroupMmaLoadMatrixOp::verify() {
2149 auto srcType = getSrcMemref().getType();
2150 auto resType = getRes().getType();
2151 auto resMatrixType = llvm::cast<gpu::MMAMatrixType>(resType);
2152 auto operand = resMatrixType.getOperand();
2153 auto srcMemrefType = llvm::cast<MemRefType>(srcType);
2155 if (!srcMemrefType.isLastDimUnitStride())
2157 "expected source memref most minor dim must have unit stride");
2159 if (operand !=
"AOp" && operand !=
"BOp" && operand !=
"COp")
2160 return emitError(
"only AOp, BOp and COp can be loaded");
2169LogicalResult SubgroupMmaStoreMatrixOp::verify() {
2170 auto srcType = getSrc().getType();
2171 auto dstType = getDstMemref().getType();
2172 auto srcMatrixType = llvm::cast<gpu::MMAMatrixType>(srcType);
2173 auto dstMemrefType = llvm::cast<MemRefType>(dstType);
2175 if (!dstMemrefType.isLastDimUnitStride())
2177 "expected destination memref most minor dim must have unit stride");
2179 if (srcMatrixType.getOperand() !=
"COp")
2181 "expected the operand matrix being stored to have 'COp' operand type");
2190LogicalResult SubgroupMmaComputeOp::verify() {
2191 enum OperandMap {
A,
B,
C };
2192 SmallVector<MMAMatrixType, 3> opTypes;
2193 opTypes.push_back(llvm::cast<MMAMatrixType>(getOpA().
getType()));
2194 opTypes.push_back(llvm::cast<MMAMatrixType>(getOpB().
getType()));
2195 opTypes.push_back(llvm::cast<MMAMatrixType>(getOpC().
getType()));
2197 if (opTypes[A].getOperand() !=
"AOp" || opTypes[B].getOperand() !=
"BOp" ||
2198 opTypes[C].getOperand() !=
"COp")
2199 return emitError(
"operands must be in the order AOp, BOp, COp");
2201 ArrayRef<int64_t> aShape, bShape, cShape;
2202 aShape = opTypes[
A].getShape();
2203 bShape = opTypes[
B].getShape();
2204 cShape = opTypes[
C].getShape();
2206 if (aShape[1] != bShape[0] || aShape[0] != cShape[0] ||
2207 bShape[1] != cShape[1])
2208 return emitError(
"operand shapes do not satisfy matmul constraints");
2213LogicalResult MemcpyOp::fold(FoldAdaptor adaptor,
2214 SmallVectorImpl<::mlir::OpFoldResult> &results) {
2218LogicalResult MemsetOp::fold(FoldAdaptor adaptor,
2219 SmallVectorImpl<::mlir::OpFoldResult> &results) {
2232struct EraseRedundantGpuWaitOpPairs :
public OpRewritePattern<WaitOp> {
2236 LogicalResult matchAndRewrite(WaitOp op,
2237 PatternRewriter &rewriter)
const final {
2238 auto predicate = [](Value value) {
2239 auto waitOp = value.getDefiningOp<WaitOp>();
2240 return waitOp && waitOp->getNumOperands() == 0;
2242 if (llvm::none_of(op.getAsyncDependencies(), predicate))
2244 SmallVector<Value> validOperands;
2245 for (Value operand : op->getOperands()) {
2246 if (predicate(operand))
2248 validOperands.push_back(operand);
2250 rewriter.
modifyOpInPlace(op, [&]() { op->setOperands(validOperands); });
2262struct SimplifyGpuWaitOp :
public OpRewritePattern<WaitOp> {
2266 LogicalResult matchAndRewrite(WaitOp op,
2267 PatternRewriter &rewriter)
const final {
2270 if (op.getAsyncDependencies().empty() && !op.getAsyncToken()) {
2275 if (llvm::hasSingleElement(op.getAsyncDependencies()) &&
2276 op.getAsyncToken()) {
2277 rewriter.
replaceOp(op, op.getAsyncDependencies());
2281 if (op.getAsyncToken() && op.getAsyncToken().use_empty()) {
2291void WaitOp::getCanonicalizationPatterns(RewritePatternSet &results,
2292 MLIRContext *context) {
2293 results.
add<EraseRedundantGpuWaitOpPairs, SimplifyGpuWaitOp>(context);
2300LogicalResult AllocOp::verify() {
2301 auto memRefType = llvm::cast<MemRefType>(getMemref().
getType());
2307 unsigned numSymbols = 0;
2308 if (!memRefType.getLayout().isIdentity())
2309 numSymbols = memRefType.getLayout().getAffineMap().getNumSymbols();
2310 if (getSymbolOperands().size() != numSymbols) {
2312 "symbol operand count does not equal memref symbol count");
2322struct SimplifyDimOfAllocOp :
public OpRewritePattern<memref::DimOp> {
2323 using OpRewritePattern<memref::DimOp>::OpRewritePattern;
2325 LogicalResult matchAndRewrite(memref::DimOp dimOp,
2326 PatternRewriter &rewriter)
const override {
2327 std::optional<int64_t> index = dimOp.getConstantIndex();
2331 int64_t indexVal = index.value();
2332 auto memrefType = llvm::dyn_cast<MemRefType>(dimOp.getSource().getType());
2333 if (!memrefType || indexVal < 0 || indexVal >= memrefType.getRank() ||
2334 !memrefType.isDynamicDim(indexVal))
2337 auto alloc = dimOp.getSource().getDefiningOp<AllocOp>();
2341 Value substituteOp = *(alloc.getDynamicSizes().begin() +
2342 memrefType.getDynamicDimIndex(indexVal));
2343 rewriter.
replaceOp(dimOp, substituteOp);
2350void AllocOp::getCanonicalizationPatterns(RewritePatternSet &results,
2351 MLIRContext *context) {
2352 results.
add<SimplifyDimOfAllocOp>(context);
2360 Attribute
target, CompilationTarget format,
2361 StringAttr
object, DictionaryAttr properties,
2362 KernelTableAttr kernels) {
2364 return emitError() <<
"the target attribute cannot be null";
2365 if (
target.hasPromiseOrImplementsInterface<TargetAttrInterface>())
2367 return emitError() <<
"the target attribute must implement or promise the "
2368 "`gpu::TargetAttrInterface`";
2372ParseResult parseObject(AsmParser &odsParser, CompilationTarget &format,
2373 StringAttr &
object) {
2374 std::optional<CompilationTarget> formatResult;
2375 StringRef enumKeyword;
2378 formatResult = CompilationTarget::Fatbin;
2379 if (!formatResult &&
2381 gpu::symbolizeEnum<gpu::CompilationTarget>(enumKeyword)) &&
2383 return odsParser.
emitError(loc,
"expected an equal sign");
2385 return odsParser.
emitError(loc,
"expected keyword for GPU object format");
2386 FailureOr<StringAttr> objectResult =
2387 FieldParser<StringAttr>::parse(odsParser);
2388 if (
failed(objectResult))
2390 "failed to parse GPU_ObjectAttr parameter "
2391 "'object' which is to be a `StringAttr`");
2392 format = *formatResult;
2393 object = *objectResult;
2397void printObject(AsmPrinter &odsParser, CompilationTarget format,
2398 StringAttr
object) {
2399 if (format != CompilationTarget::Fatbin)
2400 odsParser << stringifyEnum(format) <<
" = ";
2401 odsParser << object;
2414 if (
auto intAttr = mlir::dyn_cast<IntegerAttr>(
target)) {
2415 if (intAttr.getInt() < 0) {
2416 return emitError() <<
"the object index must be positive";
2418 }
else if (!
target.hasPromiseOrImplementsInterface<TargetAttrInterface>()) {
2420 <<
"the target attribute must be a GPU Target attribute";
2430LogicalResult gpu::DynamicSharedMemoryOp::verify() {
2431 if (!getOperation()->getParentWithTrait<OpTrait::SymbolTable>())
2432 return emitOpError() <<
"must be inside an op with symbol table";
2434 MemRefType memrefType = getResultMemref().getType();
2436 if (!GPUDialect::hasWorkgroupMemoryAddressSpace(memrefType)) {
2438 << gpu::AddressSpaceAttr::getMnemonic() <<
"<"
2439 << stringifyEnum(gpu::AddressSpace::Workgroup) <<
">";
2441 if (memrefType.hasStaticShape()) {
2442 return emitOpError() <<
"result memref type must be memref<?xi8, "
2443 "#gpu.address_space<workgroup>>";
2452void WarpExecuteOnLane0Op::print(OpAsmPrinter &p) {
2453 p <<
"(" << getLaneid() <<
")";
2455 SmallVector<StringRef> coreAttr = {getWarpSizeAttrName()};
2456 auto warpSizeAttr = getOperation()->getAttr(getWarpSizeAttrName());
2457 p <<
"[" << llvm::cast<IntegerAttr>(warpSizeAttr).getInt() <<
"]";
2459 if (!getArgs().empty())
2460 p <<
" args(" << getArgs() <<
" : " << getArgs().getTypes() <<
")";
2461 if (!getResults().empty())
2462 p <<
" -> (" << getResults().getTypes() <<
')';
2466 !getResults().empty());
2470ParseResult WarpExecuteOnLane0Op::parse(OpAsmParser &parser,
2471 OperationState &
result) {
2473 result.regions.reserve(1);
2474 Region *warpRegion =
result.addRegion();
2477 OpAsmParser::UnresolvedOperand laneId;
2489 result.addAttribute(getWarpSizeAttrName(OperationName(getOperationName(),
2496 llvm::SMLoc inputsOperandsLoc;
2497 SmallVector<OpAsmParser::UnresolvedOperand> inputsOperands;
2498 SmallVector<Type> inputTypes;
2508 if (parser.
resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc,
2519 WarpExecuteOnLane0Op::ensureTerminator(*warpRegion, builder,
result.location);
2527void WarpExecuteOnLane0Op::getSuccessorRegions(
2528 RegionBranchPoint point, SmallVectorImpl<RegionSuccessor> ®ions) {
2530 regions.push_back(RegionSuccessor(getOperation()));
2535 regions.push_back(RegionSuccessor(&getWarpRegion()));
2538ValueRange WarpExecuteOnLane0Op::getSuccessorInputs(RegionSuccessor successor) {
2541void WarpExecuteOnLane0Op::build(OpBuilder &builder, OperationState &
result,
2544 build(builder,
result, resultTypes, laneId, warpSize,
2548void WarpExecuteOnLane0Op::build(OpBuilder &builder, OperationState &
result,
2552 result.addOperands(laneId);
2553 result.addAttribute(getAttributeNames()[0],
2555 result.addTypes(resultTypes);
2556 result.addOperands(args);
2557 assert(args.size() == blockArgTypes.size());
2558 OpBuilder::InsertionGuard guard(builder);
2559 Region *warpRegion =
result.addRegion();
2561 for (
auto [type, arg] : llvm::zip_equal(blockArgTypes, args))
2570 if (expanded == distributed)
2572 auto expandedVecType = llvm::dyn_cast<VectorType>(expanded);
2573 auto distributedVecType = llvm::dyn_cast<VectorType>(distributed);
2574 if (!expandedVecType || !distributedVecType)
2575 return op->
emitOpError(
"expected vector type for distributed operands.");
2576 if (expandedVecType.getRank() != distributedVecType.getRank() ||
2577 expandedVecType.getElementType() != distributedVecType.getElementType())
2579 "expected distributed vectors to have same rank and element type.");
2582 for (
int64_t i = 0, e = expandedVecType.getRank(); i < e; i++) {
2583 int64_t eDim = expandedVecType.getDimSize(i);
2584 int64_t dDim = distributedVecType.getDimSize(i);
2587 if (eDim % dDim != 0)
2589 <<
"expected expanded vector dimension #" << i <<
" (" << eDim
2590 <<
") to be a multipler of the distributed vector dimension ("
2592 scales[i] = eDim / dDim;
2594 if (llvm::product_of(scales) != warpSize)
2596 <<
"incompatible distribution dimensions from " << expandedVecType
2597 <<
" to " << distributedVecType <<
" with warp size = " << warpSize;
2602LogicalResult WarpExecuteOnLane0Op::verify() {
2603 if (getArgs().size() != getWarpRegion().getNumArguments())
2605 "expected same number op arguments and block arguments.");
2606 auto yield = dyn_cast<gpu::YieldOp>(getBody()->getTerminator());
2608 return emitOpError(
"expected body to be terminated with 'gpu.yield'");
2609 if (yield.getNumOperands() != getNumResults())
2611 "expected same number of yield operands and return values.");
2612 int64_t warpSize = getWarpSize();
2613 for (
auto [regionArg, arg] :
2614 llvm::zip_equal(getWarpRegion().getArguments(), getArgs())) {
2616 warpSize, getOperation())))
2619 for (
auto [yieldOperand,
result] :
2620 llvm::zip_equal(yield.getOperands(), getResults())) {
2622 warpSize, getOperation())))
2627bool WarpExecuteOnLane0Op::areTypesCompatible(Type
lhs, Type
rhs) {
2632gpu::YieldOp WarpExecuteOnLane0Op::getTerminator() {
2633 return cast<gpu::YieldOp>(getBody()->getTerminator());
2640void gpu::SubgroupBroadcastOp::inferResultRanges(
2641 ArrayRef<ConstantIntRanges> argRanges,
SetIntRangeFn setResultRange) {
2642 setResultRange(getResult(), argRanges.front());
2646 switch (getBroadcastType()) {
2647 case BroadcastType::first_active_lane:
2651 case BroadcastType::specific_lane:
2655 llvm_unreachable(
"Unknown BroadcastType");
2658LogicalResult gpu::SubgroupBroadcastOp::verify() {
2659 switch (getBroadcastType()) {
2660 case BroadcastType::first_active_lane:
2663 <<
"lane can only be specified for `specific_lane` broadcast";
2665 case BroadcastType::specific_lane:
2668 <<
"lane must be specified for `specific_lane` broadcast";
2671 llvm_unreachable(
"Unknown BroadcastType");
2674OpFoldResult gpu::SubgroupBroadcastOp::fold(FoldAdaptor ) {
2676 if (
auto prev = getSrc().getDefiningOp<SubgroupBroadcastOp>())
2677 return prev.getResult();
2692KernelMetadataAttr KernelMetadataAttr::get(FunctionOpInterface kernel,
2693 DictionaryAttr metadata) {
2694 assert(kernel &&
"invalid kernel");
2695 return get(kernel.getNameAttr(), kernel.getFunctionType(),
2696 kernel.getAllArgAttrs(), metadata);
2701 FunctionOpInterface kernel,
2702 DictionaryAttr metadata) {
2703 assert(kernel &&
"invalid kernel");
2705 kernel.getAllArgAttrs(), metadata);
2709KernelMetadataAttr::appendMetadata(ArrayRef<NamedAttribute> attrs)
const {
2712 NamedAttrList attrList;
2713 if (DictionaryAttr dict = getMetadata())
2716 return KernelMetadataAttr::get(getName(), getFunctionType(),
getArgAttrs(),
2722 StringAttr name, Type functionType,
2723 ArrayAttr argAttrs, DictionaryAttr metadata) {
2725 return emitError() <<
"the kernel name can't be empty";
2727 if (llvm::any_of(argAttrs, [](Attribute attr) {
2728 return !llvm::isa<DictionaryAttr>(attr);
2731 <<
"all attributes in the array must be a dictionary attribute";
2740KernelTableAttr KernelTableAttr::get(MLIRContext *context,
2741 ArrayRef<KernelMetadataAttr> kernels,
2744 assert((!isSorted || llvm::is_sorted(kernels)) &&
2745 "expected a sorted kernel array");
2747 if (isSorted || llvm::is_sorted(kernels))
2748 return Base::get(context, kernels);
2750 SmallVector<KernelMetadataAttr> kernelsTmp(kernels);
2751 llvm::array_pod_sort(kernelsTmp.begin(), kernelsTmp.end());
2752 return Base::get(context, kernelsTmp);
2755KernelTableAttr KernelTableAttr::getChecked(
2757 ArrayRef<KernelMetadataAttr> kernels,
bool isSorted) {
2759 assert((!isSorted || llvm::is_sorted(kernels)) &&
2760 "expected a sorted kernel array");
2762 if (isSorted || llvm::is_sorted(kernels))
2763 return Base::getChecked(
emitError, context, kernels);
2765 SmallVector<KernelMetadataAttr> kernelsTmp(kernels);
2766 llvm::array_pod_sort(kernelsTmp.begin(), kernelsTmp.end());
2767 return Base::getChecked(
emitError, context, kernelsTmp);
2772 ArrayRef<KernelMetadataAttr> kernels) {
2773 if (kernels.size() < 2)
2776 if (std::adjacent_find(kernels.begin(), kernels.end(),
2777 [](KernelMetadataAttr l, KernelMetadataAttr r) {
2778 return l.getName() == r.getName();
2779 }) != kernels.end()) {
2780 return emitError() <<
"expected all kernels to be uniquely named";
2785KernelMetadataAttr KernelTableAttr::lookup(StringRef key)
const {
2787 return found ? *iterator : KernelMetadataAttr();
2790KernelMetadataAttr KernelTableAttr::lookup(StringAttr key)
const {
2792 return found ? *iterator : KernelMetadataAttr();
2872 return CompilationTarget::Fatbin;
2875std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
2877 std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
options;
2878 llvm::StringSaver stringSaver(
options.first);
2884 if (!opts.empty() && opts.front() ==
'"' && opts.back() ==
'"')
2885 opts.consume_front(
"\""), opts.consume_back(
"\"");
2886 if (!opts.empty() && opts.front() ==
'\'' && opts.back() ==
'\'')
2887 opts.consume_front(
"'"), opts.consume_back(
"'");
2889 llvm::cl::TokenizeWindowsCommandLine(opts, stringSaver,
options.second,
2892 llvm::cl::TokenizeGNUCommandLine(opts, stringSaver,
options.second,
2898std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
2903std::pair<llvm::BumpPtrAllocator, SmallVector<const char *>>
2905 size_t startPos =
cmdOptions.find(startsWith);
2906 if (startPos == std::string::npos)
2917#include "mlir/Dialect/GPU/IR/GPUOpInterfaces.cpp.inc"
2918#include "mlir/Dialect/GPU/IR/GPUOpsEnums.cpp.inc"
2920#define GET_ATTRDEF_CLASSES
2921#include "mlir/Dialect/GPU/IR/GPUOpsAttributes.cpp.inc"
2923#define GET_OP_CLASSES
2924#include "mlir/Dialect/GPU/IR/GPUOps.cpp.inc"
2926#include "mlir/Dialect/GPU/IR/CompilationAttrInterfaces.cpp.inc"
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
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 ParseResult parseAllReduceOperation(AsmParser &parser, AllReduceOperationAttr &attr)
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 void printAllReduceOperation(AsmPrinter &printer, Operation *op, AllReduceOperationAttr attr)
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)
This base class exposes generic asm parser hooks, usable across the various derived parsers.
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.
This base class exposes generic asm printer hooks, usable across the various derived printers.
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 insertOperands(unsigned index, ValueRange operands)
Insert the given operands into the operand list at the given 'index'.
AttrClass getAttrOfType(StringAttr name)
Block * getBlock()
Returns the operation block that contains this operation.
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...
void setAttr(StringAttr name, Attribute value)
If the an attribute exists with the specified name, change it to the new value.
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 StringRef getSymbolAttrName()
Return the name of the attribute used for symbol names.
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 ....