MLIR 24.0.0git
GroupOps.cpp
Go to the documentation of this file.
1//===- GroupOps.cpp - MLIR SPIR-V Group Ops ------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// Defines the group operations in the SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
15
16#include "SPIRVOpUtils.h"
17#include "SPIRVParsingUtils.h"
18
19using namespace mlir::spirv::AttrNames;
20
21namespace mlir::spirv {
22
23template <typename OpTy>
24static LogicalResult verifyGroupNonUniformArithmeticOp(Operation *groupOp) {
25 GroupOperation operation = cast<OpTy>(groupOp).getGroupOperation();
26 if (operation == GroupOperation::ClusteredReduce &&
27 groupOp->getNumOperands() == 1)
28 return groupOp->emitOpError("cluster size operand must be provided for "
29 "'ClusteredReduce' group operation");
30 if (groupOp->getNumOperands() > 1) {
31 Operation *sizeOp = groupOp->getOperand(1).getDefiningOp();
32 int32_t clusterSize = 0;
33
34 // TODO: support specialization constant here.
35 if (failed(extractValueFromConstOp(sizeOp, clusterSize)))
36 return groupOp->emitOpError(
37 "cluster size operand must come from a constant op");
38
39 if (!llvm::isPowerOf2_32(clusterSize))
40 return groupOp->emitOpError(
41 "cluster size operand must be a power of two");
42 }
43 return success();
44}
45
46//===----------------------------------------------------------------------===//
47// spirv.GroupBroadcast
48//===----------------------------------------------------------------------===//
49
50LogicalResult GroupBroadcastOp::verify() {
51 if (auto localIdTy = dyn_cast<VectorType>(getLocalid().getType()))
52 if (localIdTy.getNumElements() != 2 && localIdTy.getNumElements() != 3)
53 return emitOpError("localid is a vector and can be with only "
54 " 2 or 3 components, actual number is ")
55 << localIdTy.getNumElements();
56
57 return success();
58}
59
60//===----------------------------------------------------------------------===//
61// spirv.GroupNonUniformBroadcast
62//===----------------------------------------------------------------------===//
63
64LogicalResult GroupNonUniformBroadcastOp::verify() {
65 // SPIR-V spec: "Before version 1.5, Id must come from a
66 // constant instruction.
67 auto targetEnv = spirv::getDefaultTargetEnv(getContext());
68 if (auto spirvModule = (*this)->getParentOfType<spirv::ModuleOp>())
69 targetEnv = spirv::lookupTargetEnvOrDefault(spirvModule);
70
71 if (targetEnv.getVersion() < spirv::Version::V_1_5) {
72 auto *idOp = getId().getDefiningOp();
73 if (!idOp || !isa<spirv::ConstantOp, // for normal constant
74 spirv::ReferenceOfOp>(idOp)) // for spec constant
75 return emitOpError("id must be the result of a constant op");
76 }
77
78 return success();
79}
80
81//===----------------------------------------------------------------------===//
82// spirv.GroupNonUniformShuffle*
83//===----------------------------------------------------------------------===//
84
85template <typename OpTy>
86static LogicalResult verifyGroupNonUniformShuffleOp(OpTy op) {
87 if (op.getOperands().back().getType().isSignedInteger())
88 return op.emitOpError("second operand must be a singless/unsigned integer");
89
90 return success();
91}
92
93LogicalResult GroupNonUniformShuffleOp::verify() {
95}
96LogicalResult GroupNonUniformShuffleDownOp::verify() {
98}
99LogicalResult GroupNonUniformShuffleUpOp::verify() {
100 return verifyGroupNonUniformShuffleOp(*this);
101}
102LogicalResult GroupNonUniformShuffleXorOp::verify() {
103 return verifyGroupNonUniformShuffleOp(*this);
104}
105
106//===----------------------------------------------------------------------===//
107// spirv.GroupNonUniformFAddOp
108//===----------------------------------------------------------------------===//
109
110LogicalResult GroupNonUniformFAddOp::verify() {
112}
113
114//===----------------------------------------------------------------------===//
115// spirv.GroupNonUniformFMaxOp
116//===----------------------------------------------------------------------===//
117
118LogicalResult GroupNonUniformFMaxOp::verify() {
120}
121
122//===----------------------------------------------------------------------===//
123// spirv.GroupNonUniformFMinOp
124//===----------------------------------------------------------------------===//
125
126LogicalResult GroupNonUniformFMinOp::verify() {
128}
129
130//===----------------------------------------------------------------------===//
131// spirv.GroupNonUniformFMulOp
132//===----------------------------------------------------------------------===//
133
134LogicalResult GroupNonUniformFMulOp::verify() {
136}
137
138//===----------------------------------------------------------------------===//
139// spirv.GroupNonUniformIAddOp
140//===----------------------------------------------------------------------===//
141
142LogicalResult GroupNonUniformIAddOp::verify() {
144}
145
146//===----------------------------------------------------------------------===//
147// spirv.GroupNonUniformIMulOp
148//===----------------------------------------------------------------------===//
149
150LogicalResult GroupNonUniformIMulOp::verify() {
152}
153
154//===----------------------------------------------------------------------===//
155// spirv.GroupNonUniformSMaxOp
156//===----------------------------------------------------------------------===//
157
158LogicalResult GroupNonUniformSMaxOp::verify() {
160}
161
162//===----------------------------------------------------------------------===//
163// spirv.GroupNonUniformSMinOp
164//===----------------------------------------------------------------------===//
165
166LogicalResult GroupNonUniformSMinOp::verify() {
168}
169
170//===----------------------------------------------------------------------===//
171// spirv.GroupNonUniformUMaxOp
172//===----------------------------------------------------------------------===//
173
174LogicalResult GroupNonUniformUMaxOp::verify() {
176}
177
178//===----------------------------------------------------------------------===//
179// spirv.GroupNonUniformUMinOp
180//===----------------------------------------------------------------------===//
181
182LogicalResult GroupNonUniformUMinOp::verify() {
184}
185
186//===----------------------------------------------------------------------===//
187// spirv.GroupNonUniformBitwiseAnd
188//===----------------------------------------------------------------------===//
189
190LogicalResult GroupNonUniformBitwiseAndOp::verify() {
192}
193
194//===----------------------------------------------------------------------===//
195// spirv.GroupNonUniformBitwiseOr
196//===----------------------------------------------------------------------===//
197
198LogicalResult GroupNonUniformBitwiseOrOp::verify() {
200}
201
202//===----------------------------------------------------------------------===//
203// spirv.GroupNonUniformBitwiseXor
204//===----------------------------------------------------------------------===//
205
206LogicalResult GroupNonUniformBitwiseXorOp::verify() {
208}
209
210//===----------------------------------------------------------------------===//
211// spirv.GroupNonUniformLogicalAnd
212//===----------------------------------------------------------------------===//
213
214LogicalResult GroupNonUniformLogicalAndOp::verify() {
216}
217
218//===----------------------------------------------------------------------===//
219// spirv.GroupNonUniformLogicalOr
220//===----------------------------------------------------------------------===//
221
222LogicalResult GroupNonUniformLogicalOrOp::verify() {
224}
225
226//===----------------------------------------------------------------------===//
227// spirv.GroupNonUniformLogicalXor
228//===----------------------------------------------------------------------===//
229
230LogicalResult GroupNonUniformLogicalXorOp::verify() {
232}
233
234//===----------------------------------------------------------------------===//
235// spirv.GroupNonUniformRotateKHR
236//===----------------------------------------------------------------------===//
237
238LogicalResult GroupNonUniformRotateKHROp::verify() {
239 if (Value clusterSizeVal = getClusterSize()) {
240 mlir::Operation *defOp = clusterSizeVal.getDefiningOp();
241 int32_t clusterSize = 0;
242
243 if (failed(extractValueFromConstOp(defOp, clusterSize)))
244 return emitOpError("cluster size operand must come from a constant op");
245
246 if (!llvm::isPowerOf2_32(clusterSize))
247 return emitOpError("cluster size operand must be a power of two");
248 }
249
250 return success();
251}
252
253} // namespace mlir::spirv
return success()
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.
b getContext())
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
unsigned getNumOperands()
Definition Operation.h:371
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
static LogicalResult verifyGroupNonUniformArithmeticOp(Operation *groupOp)
Definition GroupOps.cpp:24
TargetEnvAttr lookupTargetEnvOrDefault(Operation *op)
Queries the target environment recursively from enclosing symbol table ops containing the given op or...
static LogicalResult verifyGroupNonUniformShuffleOp(OpTy op)
Definition GroupOps.cpp:86
LogicalResult extractValueFromConstOp(Operation *op, int32_t &value)
Definition SPIRVOps.cpp:49
TargetEnvAttr getDefaultTargetEnv(MLIRContext *context)
Returns the default target environment: SPIR-V 1.0 with Shader capability and no extra extensions.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307