MLIR 24.0.0git
TosaValidation.cpp
Go to the documentation of this file.
1//===- TosaValidation.cpp ------------------------------------------------===//
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// Validate if TOSA dialect input matches with the specification for given
10// requirements.
11//
12//===----------------------------------------------------------------------===//
13
17
18#include <string>
19#include <type_traits>
20
24#include "mlir/IR/Builders.h"
25#include "mlir/IR/BuiltinOps.h"
26#include "mlir/IR/Matchers.h"
28#include "mlir/Pass/Pass.h"
30#include "llvm/ADT/STLExtras.h"
31#include "llvm/ADT/StringExtras.h"
32#include "llvm/ADT/TypeSwitch.h"
33#include "llvm/Support/FormatVariadic.h"
34
35namespace mlir {
36namespace tosa {
37#define GEN_PASS_DEF_TOSAVALIDATION
38#include "mlir/Dialect/Tosa/Transforms/Passes.h.inc"
39} // namespace tosa
40} // namespace mlir
41
42using namespace mlir;
43using namespace mlir::tosa;
44
45namespace {
46
47static LogicalResult
48checkConstantOperands(Operation *op, ArrayRef<unsigned int> operandIndices) {
49 for (const auto index : operandIndices) {
50 Attribute attr;
51 if (!matchPattern(op->getOperand(index), m_Constant(&attr))) {
52 return op->emitOpError("expected compile time resolvable constant, but "
53 "got variable value for operand #")
54 << index;
55 }
56 }
57 return success();
58}
59
60static LogicalResult checkConstantOperandMul(Operation *op,
61 const TargetEnv &env) {
62 if (!env.allows(Extension::dynamic) && isa<tosa::MulOp>(op)) {
63 // Check 'shift'
64 return checkConstantOperands(op, {2});
65 }
66 return success();
67}
68
69static LogicalResult checkConstantOperandTable(Operation *op,
70 const TargetEnv &env) {
71 if (!env.allows(Extension::dynamic) && isa<tosa::TableOp>(op)) {
72 // Check 'table'
73 return checkConstantOperands(op, {1});
74 }
75 return success();
76}
77
78static LogicalResult checkConstantOperandPad(Operation *op,
79 const TargetEnv &env) {
80 if (auto padOp = dyn_cast<tosa::PadOp>(op)) {
81 // Assume this op is zero-padding if padConst is not presented
82 if (!env.allows(Extension::dynamic) && padOp.getPadConst())
83 // Check 'pad_const'
84 // Note: 'padding' (operand 1) is not checked as it is a tosa.shape type
85 return checkConstantOperands(op, {2});
86 }
87 return success();
88}
89
90static LogicalResult checkConstantOperandRescale(Operation *op,
91 const TargetEnv &env) {
92 if (!env.allows(Extension::dynamic) && isa<tosa::RescaleOp>(op)) {
93 // Check 'multiplier', 'shift', 'input_zp' and 'output_zp'
94 return checkConstantOperands(op, {1, 2, 3, 4});
95 }
96 return success();
97}
98
99template <typename T>
100static LogicalResult checkConstantOperandConvOps(Operation *op,
101 const TargetEnv &env) {
102 if (!env.allows(Extension::dynamic) && isa<T>(op)) {
103 // Check 'input_zp' and 'weight_zp'
104 return checkConstantOperands(op, {3, 4});
105 }
106 return success();
107}
108
109static LogicalResult checkConstantOperandMatMul(Operation *op,
110 const TargetEnv &env) {
111 if (!env.allows(Extension::dynamic) &&
112 isa<tosa::MatMulOp, tosa::MatMulTOp>(op)) {
113 // Check 'A_zp' and 'B_zp'
114 return checkConstantOperands(op, {2, 3});
115 }
116 return success();
117}
118
119static LogicalResult
120checkConstantOperandRowGatherBlockScaled(Operation *op, const TargetEnv &env) {
121 if (!env.allows(Extension::dynamic) &&
122 isa<tosa::RowGatherBlockScaledOp>(op)) {
123 auto rowGatherOp = cast<tosa::RowGatherBlockScaledOp>(op);
124 const unsigned rowCountIndex = rowGatherOp.getValues().size() + 1;
125 return checkConstantOperands(op, {rowCountIndex});
126 }
127 return success();
128}
129
130static LogicalResult checkConstantOperandRowGather(Operation *op,
131 const TargetEnv &env) {
132 if (!env.allows(Extension::dynamic) && isa<tosa::RowGatherOp>(op)) {
133 // Check 'row_count'
134 return checkConstantOperands(op, {2});
135 }
136 return success();
137}
138
139static LogicalResult checkConstantOperandAvgPool2d(Operation *op,
140 const TargetEnv &env) {
141 if (!env.allows(Extension::dynamic) && isa<tosa::AvgPool2dOp>(op)) {
142 // Check 'input_zp' and 'output_zp'
143 return checkConstantOperands(op, {1, 2});
144 }
145 return success();
146}
147
148static LogicalResult
149checkConstantOperandAvgPool2dAdaptive(Operation *op, const TargetEnv &env) {
150 if (!env.allows(Extension::dynamic) && isa<tosa::AvgPool2dAdaptiveOp>(op)) {
151 // Check 'input_zp' and 'output_zp'.
152 // Note: 'kernel', 'stride', and 'pad' (operands 3, 4, 5) are not checked
153 // as they are tosa.shape types.
154 return checkConstantOperands(op, {1, 2});
155 }
156 return success();
157}
158
159static LogicalResult checkConstantOperandNegate(Operation *op,
160 const TargetEnv &env) {
161 if (!env.allows(Extension::dynamic) && isa<tosa::NegateOp>(op)) {
162 // Check 'input1_zp' and 'output_zp'
163 return checkConstantOperands(op, {1, 2});
164 }
165 return success();
166}
167
168static LogicalResult checkConstantOperandSilceShape(Operation *op,
169 const TargetEnv &env) {
170 if (!env.allows(Extension::dynamic) && isa<tosa::SliceShapeOp>(op)) {
171 // Check 'start' and 'size'
172 return checkConstantOperands(op, {1, 2});
173 }
174 return success();
175}
176
177// MATMUL's data type availability predates variadic batch support, so the
178// generated availability checks cannot distinguish its 1.0 and 1.1 shapes.
179static LogicalResult checkSpecificationVersionConstraint(Operation *op,
180 const TargetEnv &env) {
181 auto matmul = dyn_cast<tosa::MatMulOp>(op);
182 if (!matmul ||
184 TosaSpecificationVersion(SpecificationVersion::V_1_1_DRAFT)))
185 return success();
186
187 auto aType = dyn_cast<RankedTensorType>(matmul.getA().getType());
188 auto bType = dyn_cast<RankedTensorType>(matmul.getB().getType());
189 auto outputType = dyn_cast<RankedTensorType>(matmul.getOutput().getType());
190 auto getBatchDimOrDynamic = [](RankedTensorType type) {
191 return type ? type.getDimSize(0) : ShapedType::kDynamic;
192 };
193 if ((!aType || aType.getRank() == 3) && (!bType || bType.getRank() == 3) &&
194 (!outputType || outputType.getRank() == 3) &&
195 succeeded(verifyCompatibleDims({getBatchDimOrDynamic(aType),
196 getBatchDimOrDynamic(bType),
197 getBatchDimOrDynamic(outputType)})))
198 return success();
199
200 return op->emitOpError(
201 "MATMUL ranks other than 3 or batch broadcasting require TOSA "
202 "specification version 1.1.draft");
203}
204
205//===----------------------------------------------------------------------===//
206// TOSA Validation Pass.
207//===----------------------------------------------------------------------===//
208
209struct TosaValidation : public tosa::impl::TosaValidationBase<TosaValidation> {
210public:
211 explicit TosaValidation() { populateConstantOperandChecks(); }
212
213 explicit TosaValidation(const TosaValidationOptions &options)
214 : TosaValidation() {
215 this->strictOpSpecAlignment = options.strictOpSpecAlignment;
216 this->allowInvalidOpDatatypeCombinations =
217 options.allowInvalidOpDatatypeCombinations;
218 this->validateFunctionSignature = options.validateFunctionSignature;
219 }
220 void runOnOperation() final;
221
222 LogicalResult applyConstantOperandCheck(Operation *op) {
223 for (auto &checker : constCheckers) {
224 if (failed(checker(op, targetEnv)))
225 return failure();
226 }
227 return success();
228 }
229
230 LogicalResult applyFunctionSignatureCheck(func::FuncOp op);
231 LogicalResult applyLevelCheck(Operation *op);
232 LogicalResult applyAttributeCheck(Operation *op);
233
234 // check variable read/write data types against variable declarations
235 LogicalResult applyVariableCheck(Operation *op);
236
237 // check error if conditions
238 LogicalResult applyErrorIfCheck(Operation *op);
239
240private:
241 void populateConstantOperandChecks() {
242 constCheckers.emplace_back(checkConstantOperandMul);
243 constCheckers.emplace_back(checkConstantOperandTable);
244 constCheckers.emplace_back(checkConstantOperandPad);
245 constCheckers.emplace_back(checkConstantOperandRescale);
246 constCheckers.emplace_back(checkConstantOperandConvOps<tosa::Conv2DOp>);
247 constCheckers.emplace_back(checkConstantOperandConvOps<tosa::Conv3DOp>);
248 constCheckers.emplace_back(
249 checkConstantOperandConvOps<tosa::DepthwiseConv2DOp>);
250 constCheckers.emplace_back(
251 checkConstantOperandConvOps<tosa::TransposeConv2DOp>);
252 constCheckers.emplace_back(checkConstantOperandMatMul);
253 constCheckers.emplace_back(checkConstantOperandRowGather);
254 constCheckers.emplace_back(checkConstantOperandRowGatherBlockScaled);
255 constCheckers.emplace_back(checkConstantOperandAvgPool2d);
256 constCheckers.emplace_back(checkConstantOperandAvgPool2dAdaptive);
257 constCheckers.emplace_back(checkConstantOperandNegate);
258 constCheckers.emplace_back(checkConstantOperandSilceShape);
259 }
260
261 LogicalResult levelCheck(Operation *op, const int32_t calculatedValue,
262 const int32_t maxLevel, const StringRef inputName,
263 const StringRef levelName) {
264 if (calculatedValue > maxLevel)
265 return op->emitOpError()
266 << "failed level check: " << inputName << " <= " << levelName
267 << " (" << maxLevel << "), got " << calculatedValue;
268 return success();
269 }
270
271 LogicalResult levelCheckKernel(Operation *op, int32_t v,
272 const StringRef inputName) {
273 return levelCheck(op, v, targetEnv.getLevel().MAX_KERNEL, inputName,
274 "MAX_KERNEL");
275 }
276
277 LogicalResult levelCheckStride(Operation *op, int32_t v,
278 const StringRef inputName) {
279 return levelCheck(op, v, targetEnv.getLevel().MAX_STRIDE, inputName,
280 "MAX_STRIDE");
281 }
282
283 LogicalResult levelCheckScale(Operation *op, int32_t v,
284 const StringRef inputName) {
285 return levelCheck(op, v, targetEnv.getLevel().MAX_SCALE, inputName,
286 "MAX_SCALE");
287 }
288
289 LogicalResult levelCheckListSize(Operation *op, int32_t v,
290 const StringRef inputName) {
291 const std::string inputDesc =
292 llvm::formatv("length(tensor_list_shape({0}))", inputName);
293 return levelCheck(op, v, targetEnv.getLevel().MAX_TENSOR_LIST_SIZE,
294 inputDesc, "MAX_TENSOR_LIST_SIZE");
295 }
296
297 // Perform the Level Rank check on the tensor type.
298 LogicalResult levelCheckRank(Operation *op, const Type typeToCheck,
299 const StringRef operandOrResult,
300 int32_t highest_rank) {
301 if (ShapedType type = dyn_cast<ShapedType>(typeToCheck)) {
302 if (!type.hasRank())
303 return op->emitOpError() << "failed level check: unranked tensor";
304 if (type.getRank() > highest_rank)
305 return op->emitOpError() << "failed level check: " << operandOrResult
306 << " rank(shape) <= MAX_RANK";
307 }
308 return success();
309 }
310
311 // Perform the Level Rank check on the tensor value.
312 LogicalResult levelCheckRank(Operation *op, const Value &v,
313 const StringRef operandOrResult,
314 int32_t highest_rank) {
315 return levelCheckRank(op, v.getType(), operandOrResult, highest_rank);
316 }
317
318 // Perform the Level tensor size check on the tensor type.
319 LogicalResult levelCheckSize(Operation *op, const Type &typeToCheck,
320 const StringRef operandOrResult);
321
322 // Perform the Level tensor size check on the tensor value.
323 LogicalResult levelCheckSize(Operation *op, const Value &v,
324 const StringRef operandOrResult) {
325 return levelCheckSize(op, v.getType(), operandOrResult);
326 }
327
328 // Perform the Level shape length check on a value.
329 LogicalResult levelCheckShapeLength(Operation *op, const Type typeToCheck,
330 const StringRef operandOrResult) {
331 if (tosa::shapeType shapeType = dyn_cast<tosa::shapeType>(typeToCheck)) {
332 if (shapeType.getRank() > targetEnv.getLevel().MAX_SHAPE_LEN)
333 return op->emitOpError()
334 << "failed shape type level check: " << typeToCheck
335 << " exceeds MAX_SHAPE_LEN";
336 }
337 return success();
338 }
339
340 // Level check sizes of all operands and results of the operation.
341 template <typename T>
342 LogicalResult levelCheckSizes(T tosaOp) {
343 auto op = tosaOp.getOperation();
344 for (auto v : op->getOperands()) {
345 if (failed(levelCheckSize(op, v, "operand")))
346 return failure();
347 }
348
349 for (auto v : op->getResults()) {
350 if (failed(levelCheckSize(op, v, "result")))
351 return failure();
352 }
353 return success();
354 }
355
356 // Level check ranks of all operands, attribute and results of the operation.
357 template <typename T>
358 LogicalResult levelCheckRanks(T tosaOp) {
359 auto op = tosaOp.getOperation();
360 const TosaLevel tosaLevel = targetEnv.getLevel();
361 for (auto v : op->getOperands()) {
362 if (failed(levelCheckRank(op, v, "operand", tosaLevel.MAX_RANK)))
363 return failure();
364 }
365
366 for (auto v : op->getResults()) {
367 if (failed(levelCheckRank(op, v, "result", tosaLevel.MAX_RANK)))
368 return failure();
369 }
370 return success();
371 }
372 // Level check shape lengths of all operands and results of an operation that
373 // are tosa.shape type.
374 template <typename T>
375 LogicalResult levelCheckShapeLengths(T tosaOp) {
376 for (const auto &v : tosaOp->getOperands()) {
377 if (failed(levelCheckShapeLength(tosaOp, v.getType(), "operand")))
378 return failure();
379 }
380 for (const auto &v : tosaOp->getResults()) {
381 if (failed(levelCheckShapeLength(tosaOp, v.getType(), "result")))
382 return failure();
383 }
384
385 return success();
386 }
387
388 // Level check ranks and sizes.
389 LogicalResult levelCheckRanksAndSizes(Operation *op);
390
391 // Pool Op: level check kernel/stride/pad values
392 template <typename T>
393 LogicalResult levelCheckPool(Operation *op) {
394 if (auto poolOp = dyn_cast<T>(op)) {
395 for (auto k : poolOp.getKernel()) {
396 if (failed(levelCheckKernel(op, k, "kernel"))) {
397 return failure();
398 }
399 }
400 for (auto s : poolOp.getStride()) {
401 if (failed(levelCheckStride(op, s, "stride"))) {
402 return failure();
403 }
404 }
405 for (auto p : poolOp.getPad()) {
406 if (failed(levelCheckKernel(op, p, "pad"))) {
407 return failure();
408 }
409 }
410 }
411 return success();
412 }
413
414 template <typename T>
415 static constexpr bool IsSupportedAdaptivePoolOp =
416 std::is_same_v<T, tosa::AvgPool2dAdaptiveOp> ||
417 std::is_same_v<T, tosa::MaxPool2dAdaptiveOp>;
418
419 template <typename T, typename std::enable_if<IsSupportedAdaptivePoolOp<T>,
420 int>::type = 0>
421 LogicalResult levelCheckAdaptivePool(Operation *op) {
422 auto poolOp = dyn_cast<T>(op);
423 if (!poolOp)
424 return success();
425
426 SmallVector<int64_t> kernelValues;
427 if (tosa::getConstShapeValues(poolOp.getKernel().getDefiningOp(),
428 kernelValues)) {
429 for (const auto k : kernelValues)
430 if (failed(levelCheckKernel(op, k, "kernel")))
431 return failure();
432 }
433
434 SmallVector<int64_t> strideValues;
435 if (tosa::getConstShapeValues(poolOp.getStride().getDefiningOp(),
436 strideValues)) {
437 for (const auto s : strideValues)
438 if (failed(levelCheckStride(op, s, "stride")))
439 return failure();
440 }
441
442 SmallVector<int64_t> padValues;
443 if (tosa::getConstShapeValues(poolOp.getPad().getDefiningOp(), padValues)) {
444 for (const auto p : padValues)
445 if (failed(levelCheckKernel(op, p, "pad")))
446 return failure();
447 }
448
449 return success();
450 }
451
452 // Conv Op: level check dilation/stride/pad values
453 template <typename T>
454 LogicalResult levelCheckConv(Operation *op) {
455 if (auto convOp = dyn_cast<T>(op)) {
456
457 for (auto k : convOp.getDilation()) {
458 if (failed(levelCheckKernel(op, k, "dilation"))) {
459 return failure();
460 }
461 }
462 for (auto p : convOp.getPad()) {
463 if (failed(levelCheckKernel(op, p, "pad"))) {
464 return failure();
465 }
466 }
467 for (auto s : convOp.getStride()) {
468 if (failed(levelCheckStride(op, s, "stride"))) {
469 return failure();
470 }
471 }
472 auto dilation = convOp.getDilation();
473 if (ShapedType weightType =
474 dyn_cast<ShapedType>(op->getOperand(1).getType())) {
475 auto shape = weightType.getShape();
476 if (isa<tosa::Conv2DOp>(op)) {
477 assert(shape.size() == 4);
478 assert(dilation.size() == 2);
479 if (failed(levelCheckKernel(op, dilation[0] * shape[1],
480 "dilation_y * KH")) ||
481 failed(levelCheckKernel(op, dilation[1] * shape[2],
482 "dilation_x * KW")))
483 return failure();
484 } else if (isa<tosa::Conv3DOp>(op)) {
485 assert(shape.size() == 5);
486 assert(dilation.size() == 3);
487 if (failed(levelCheckKernel(op, dilation[0] * shape[1],
488 "dilation_d * KD")) ||
489 failed(levelCheckKernel(op, dilation[1] * shape[2],
490 "dilation_y * KH")) ||
491 failed(levelCheckKernel(op, dilation[2] * shape[3],
492 "dilation_x * KW")))
493 return failure();
494 } else if (isa<tosa::DepthwiseConv2DOp>(op)) {
495 assert(shape.size() == 4);
496 assert(dilation.size() == 2);
497 if (failed(levelCheckKernel(op, dilation[0] * shape[0],
498 "dilation_y * KH")) ||
499 failed(levelCheckKernel(op, dilation[1] * shape[1],
500 "dilation_x * KW")))
501 return failure();
502 }
503 }
504 }
505 return success();
506 }
507
508 LogicalResult levelCheckConv2DBlockScaled(Operation *op) {
509 auto convOp = dyn_cast<Conv2DBlockScaledOp>(op);
510 if (!convOp)
511 return success();
512
513 SmallVector<int64_t> padValues;
514 if (tosa::getConstShapeValues(convOp.getPad().getDefiningOp(), padValues)) {
515 for (const auto p : padValues)
516 if (failed(levelCheckKernel(op, p, "pad <= MAX_KERNEL")))
517 return failure();
518 }
519
520 SmallVector<int64_t> strideValues;
521 if (tosa::getConstShapeValues(convOp.getStride().getDefiningOp(),
522 strideValues)) {
523 for (const auto s : strideValues)
524 if (failed(levelCheckKernel(op, s, "stride <= MAX_KERNEL")))
525 return failure();
526 }
527
528 SmallVector<int64_t> dilationValues;
529 if (tosa::getConstShapeValues(convOp.getDilation().getDefiningOp(),
530 dilationValues)) {
531 int64_t KH = ShapedType::kDynamic;
532 int64_t KW = ShapedType::kDynamic;
533 const ShapeAdaptor weightDataShape(convOp.getWeightData().getType());
534 KH = weightDataShape.getDimSize(1);
535 KW = weightDataShape.getDimSize(2);
536 const ShapeAdaptor weightScaleShape(convOp.getWeightScale().getType());
537 KH = ShapedType::isDynamic(KH) ? weightScaleShape.getDimSize(1) : KH;
538 KW = ShapedType::isDynamic(KW) ? weightScaleShape.getDimSize(2) : KW;
539
540 if (!ShapedType::isDynamic(KH) &&
541 failed(levelCheckKernel(op, dilationValues[0] * KH,
542 "dilation_y * KH <= MAX_KERNEL)")))
543 return failure();
544
545 if (!ShapedType::isDynamic(KW) &&
546 failed(levelCheckKernel(op, dilationValues[1] * KW,
547 "dilation_x * KW <= MAX_KERNEL)")))
548 return failure();
549 }
550
551 return success();
552 }
553
554 // FFT op: level check H, W in input shape [N,H,W]
555 template <typename T>
556 LogicalResult levelCheckFFT(Operation *op) {
557 if (isa<T>(op)) {
558 for (auto v : op->getOperands()) {
559 if (ShapedType type = dyn_cast<ShapedType>(v.getType())) {
560 auto shape = type.getShape();
561 assert(shape.size() == 3);
562 if (failed(levelCheckKernel(op, shape[1], "H")) ||
563 failed(levelCheckKernel(op, shape[2], "W"))) {
564 return failure();
565 }
566 }
567 }
568 }
569 return success();
570 }
571
572 // TransposeConv2d op: level check kH/kW, outpad, and stride
573 LogicalResult levelCheckTransposeConv2d(Operation *op) {
574 if (auto transpose = dyn_cast<tosa::TransposeConv2DOp>(op)) {
575 if (ShapedType filterType =
576 dyn_cast<ShapedType>(transpose.getWeight().getType())) {
577 auto shape = filterType.getShape();
578 assert(shape.size() == 4);
579 // level check kernel sizes for kH and KW
580 if (failed(levelCheckKernel(op, shape[1], "KH")) ||
581 failed(levelCheckKernel(op, shape[2], "KW"))) {
582 return failure();
583 }
584 }
585 for (auto p : transpose.getOutPad()) {
586 if (failed(levelCheckKernel(op, p, "pad"))) {
587 return failure();
588 }
589 }
590 for (auto s : transpose.getStride()) {
591 if (failed(levelCheckStride(op, s, "stride"))) {
592 return failure();
593 }
594 }
595 }
596 return success();
597 }
598
599 // Resize op: level check max scales
600 LogicalResult levelCheckResize(Operation *op) {
601 if (auto resize = dyn_cast<tosa::ResizeOp>(op)) {
602 SmallVector<int64_t> scale;
603 if (!tosa::getConstShapeValues(resize.getScale().getDefiningOp(),
604 scale)) {
605 return failure();
606 }
607 const int64_t scaleYN = scale[0];
608 const int64_t scaleYD = scale[1];
609 const int64_t scaleXN = scale[2];
610 const int64_t scaleXD = scale[3];
611 if (failed(
612 levelCheckScale(op, scaleYN / scaleYD, "scale_y_n/scale_y_d")) ||
613 failed(
614 levelCheckScale(op, scaleXN / scaleXD, "scale_x_n/scale_x_d"))) {
615 return failure();
616 }
617 }
618 return success();
619 }
620
621 // Recursively perform a bottom-up search to determine the maximum nesting
622 // depth, starting from a specific operation and continuing up to the function
623 // or module scope. Tosa nesting_depth starts at 0 and increments by one each
624 // time a new nested `region` is encountered.
625 static void getMaxNestedDepth(Operation *op, int32_t &depth) {
626 if (isa<mlir::func::FuncOp>(op) || isa<ModuleOp>(op))
627 return;
628
629 op = op->getParentOp();
630 if (!op)
631 return;
632
633 depth++;
634 getMaxNestedDepth(op, depth);
635 }
636
637 LogicalResult levelCheckMaxNesting(Operation *op) {
638 int32_t maxNestedDepth = 0;
639 getMaxNestedDepth(op, maxNestedDepth);
640
641 const int32_t maxNestingLevel = targetEnv.getLevel().MAX_NESTING;
642 if (maxNestedDepth >= maxNestingLevel)
643 return op->emitOpError()
644 << "failed level check: tosa_nesting_depth < MAX_NESTING" << " ("
645 << maxNestingLevel << "), got " << maxNestedDepth;
646 return success();
647 }
648
649 LogicalResult levelCheckListSize(Operation *op) {
650 if (auto concat = dyn_cast<tosa::ConcatOp>(op)) {
651 return levelCheckListSize(op, concat.getInput1().size(), "input1");
652 }
653 if (auto custom = dyn_cast<tosa::CustomOp>(op)) {
654 if (failed(levelCheckListSize(op, custom.getInputList().size(),
655 "input_list")) ||
656 failed(levelCheckListSize(op, custom.getOutputList().size(),
657 "output_list"))) {
658 return failure();
659 }
660 }
661 if (auto condIf = dyn_cast<tosa::IfOp>(op)) {
662 if (failed(
663 levelCheckListSize(op, condIf.getInputList().size(), "inputs")) ||
664 failed(levelCheckListSize(op, condIf.getOutputList().size(),
665 "outputs"))) {
666 return failure();
667 }
668 }
669 if (auto w = dyn_cast<tosa::WhileOp>(op)) {
670 if (failed(levelCheckListSize(op, w.getInputList().size(), "inputs")) ||
671 failed(levelCheckListSize(op, w.getOutputList().size(), "outputs"))) {
672 return failure();
673 }
674 }
675 if (auto concat_shape = dyn_cast<tosa::ConcatShapeOp>(op))
676 return levelCheckListSize(op, concat_shape.getInput().size(), "input");
677 return success();
678 }
679
680 LogicalResult attributeCheckRescale(Operation *op) {
681 if (auto rescale = dyn_cast<tosa::RescaleOp>(op)) {
682 if (rescale.getRoundingMode() == RoundingMode::DOUBLE_ROUND &&
683 !targetEnv.allows(Extension::doubleround)) {
684 op->emitOpError()
685 << "failed attribute check: rounding_mode = DOUBLE_ROUND "
686 << "requires extension [doubleround]";
687 return failure();
688 }
689 if (rescale.getRoundingMode() == RoundingMode::INEXACT_ROUND &&
690 !targetEnv.allows(Extension::inexactround)) {
691 op->emitOpError()
692 << "failed attribute check: rounding_mode = INEXACT_ROUND "
693 << "requires extension [inexactround]";
694 return failure();
695 }
696 }
697 return success();
698 }
699
700 LogicalResult attributeCheckCast(Operation *op) {
701 if (auto cast = dyn_cast<tosa::CastOp>(op)) {
702 const TosaSpecificationVersion targetVersion = targetEnv.getSpecVersion();
703 const TosaSpecificationVersion minRequiredVersion(1, 1, true);
704 if (cast.getInputUnsigned() &&
705 !(targetVersion.isBackwardsCompatibleWith(minRequiredVersion)))
706 return op->emitOpError()
707 << "failed attribute check: CAST attribute input_unsigned "
708 << "requires version 1.1.draft"
709 << " (got " << stringifyVersion(targetVersion) << ") ";
710 }
711 return success();
712 }
713
714 LogicalResult CheckVariable(Operation *op);
715 LogicalResult CheckVariableReadOrWrite(Operation *op);
716 LogicalResult validateValidElementType(Operation *op, Type type,
717 bool allowUnsigned = false);
718 LogicalResult validateOperationElementTypes(TosaOp op,
719 bool allowUnsigned = false);
720 LogicalResult validateOperationElementTypes(func::FuncOp op,
721 bool allowUnsigned = false);
722
723 SmallVector<
724 std::function<LogicalResult(Operation *, const tosa::TargetEnv &)>>
725 constCheckers;
727 TosaProfileCompliance profileComp;
728 tosa::TargetEnv targetEnv;
729};
730
731template <>
732LogicalResult TosaValidation::levelCheckRanks(tosa::ArgMaxOp tosaOp) {
733 auto *op = tosaOp.getOperation();
734 if (failed(levelCheckRank(op, tosaOp.getInput(), "operand",
735 targetEnv.getLevel().MAX_RANK)))
736 return failure();
737
738 // rank(output) = rank(input) - 1
739 if (failed(levelCheckRank(op, tosaOp.getOutput(), "result",
740 targetEnv.getLevel().MAX_RANK - 1)))
741 return failure();
742
743 return success();
744}
745
746template <>
747LogicalResult TosaValidation::levelCheckRanks(tosa::ArgMinOp tosaOp) {
748 auto *op = tosaOp.getOperation();
749 if (failed(levelCheckRank(op, tosaOp.getInput(), "operand",
750 targetEnv.getLevel().MAX_RANK)))
751 return failure();
752
753 // rank(output) = rank(input) - 1
754 if (failed(levelCheckRank(op, tosaOp.getOutput(), "result",
755 targetEnv.getLevel().MAX_RANK - 1)))
756 return failure();
757
758 return success();
759}
760
761template <>
762LogicalResult TosaValidation::levelCheckRanks(tosa::IfOp tosaOp) {
763 auto *op = tosaOp.getOperation();
764
765 // Only the condition input has rank limitation.
766 if (failed(levelCheckRank(op, tosaOp.getCondition(), "operand",
767 targetEnv.getLevel().MAX_RANK)))
768 return failure();
769
770 return success();
771}
772
773template <>
774LogicalResult TosaValidation::levelCheckRanks(tosa::VariableOp tosaOp) {
775 auto *op = tosaOp.getOperation();
776 auto variableType = getVariableType(tosaOp);
777 if (failed(levelCheckRank(op, variableType, "variable type",
778 targetEnv.getLevel().MAX_RANK)))
779 return failure();
780
781 return success();
782}
783
784template <>
785LogicalResult TosaValidation::levelCheckSizes(tosa::VariableOp tosaOp) {
786 auto *op = tosaOp.getOperation();
787 auto variableType = getVariableType(tosaOp);
788 if (failed(levelCheckSize(op, variableType, "variable type")))
789 return failure();
790
791 return success();
792}
793
794LogicalResult TosaValidation::levelCheckRanksAndSizes(Operation *op) {
795#define CHECK_RANKS_AND_SIZES(tosaOp) \
796 if (isa<tosa::tosaOp##Op>(op)) { \
797 if (failed(levelCheckRanks(cast<tosa::tosaOp##Op>(op)))) \
798 return failure(); \
799 if (failed(levelCheckSizes(cast<tosa::tosaOp##Op>(op)))) \
800 return failure(); \
801 }
802
803#define CHECK_SIZES(tosaOp) \
804 if (isa<tosa::tosaOp##Op>(op)) { \
805 if (failed(levelCheckSizes(cast<tosa::tosaOp##Op>(op)))) \
806 return failure(); \
807 }
808
809#define CHECK_SHAPE_LEN(tosaOp) \
810 if (isa<tosa::tosaOp##Op>(op)) { \
811 if (failed(levelCheckShapeLengths(cast<tosa::tosaOp##Op>(op)))) \
812 return failure(); \
813 }
814
815 // Tensor Operators
816 CHECK_RANKS_AND_SIZES(ArgMax);
817 CHECK_RANKS_AND_SIZES(ArgMin);
818 // Activation Functions
821 CHECK_RANKS_AND_SIZES(Sigmoid);
823 // Elementwise Binary Operators
825 CHECK_RANKS_AND_SIZES(ArithmeticRightShift);
826 CHECK_RANKS_AND_SIZES(BitwiseAnd);
827 CHECK_RANKS_AND_SIZES(BitwiseOr);
828 CHECK_RANKS_AND_SIZES(BitwiseXor);
829 CHECK_RANKS_AND_SIZES(IntDiv);
830 CHECK_RANKS_AND_SIZES(LogicalAnd);
831 CHECK_RANKS_AND_SIZES(LogicalLeftShift);
832 CHECK_RANKS_AND_SIZES(LogicalRightShift);
833 CHECK_RANKS_AND_SIZES(LogicalOr);
834 CHECK_RANKS_AND_SIZES(LogicalXor);
835 CHECK_RANKS_AND_SIZES(Maximum);
836 CHECK_RANKS_AND_SIZES(Minimum);
841 // Elementwise Unary Operators
843 CHECK_RANKS_AND_SIZES(BitwiseNot);
850 CHECK_RANKS_AND_SIZES(LogicalNot);
851 CHECK_RANKS_AND_SIZES(Negate);
852 CHECK_RANKS_AND_SIZES(Reciprocal);
855 // Elementwise Ternary Operators
856 CHECK_RANKS_AND_SIZES(Select);
857 // Comparison Operators
859 CHECK_RANKS_AND_SIZES(Greater);
860 CHECK_RANKS_AND_SIZES(GreaterEqual);
861 // Reduction Operators
862 CHECK_RANKS_AND_SIZES(ReduceAll);
863 CHECK_RANKS_AND_SIZES(ReduceAny);
864 CHECK_RANKS_AND_SIZES(ReduceMax);
865 CHECK_RANKS_AND_SIZES(ReduceMin);
866 CHECK_RANKS_AND_SIZES(ReduceProduct);
867 CHECK_RANKS_AND_SIZES(ReduceSum);
868 // Data Layout Operators
869 CHECK_RANKS_AND_SIZES(Concat);
871 CHECK_RANKS_AND_SIZES(Reshape);
872 CHECK_RANKS_AND_SIZES(ReshapeBlockScaled);
873 CHECK_RANKS_AND_SIZES(Reverse);
876 CHECK_RANKS_AND_SIZES(Transpose);
877 // Type Conversion
879 CHECK_RANKS_AND_SIZES(CastFromBlockScaled);
880 CHECK_RANKS_AND_SIZES(CastToBlockScaled);
881 CHECK_RANKS_AND_SIZES(Rescale);
882 // Data Nodes
884 CHECK_RANKS_AND_SIZES(Identity);
885 // Control Flow Operators
887 // Variable Operators
888 CHECK_RANKS_AND_SIZES(Variable);
889 CHECK_RANKS_AND_SIZES(VariableWrite);
890 CHECK_RANKS_AND_SIZES(VariableRead);
891 // Shape Operators
893
894 // For the following operators, check whether the size of each tensor
895 // operand is valid in a given Level.
896
897 // Tensor Operators
898 CHECK_SIZES(AvgPool2d);
899 CHECK_SIZES(AvgPool2dAdaptive);
900 CHECK_SIZES(Conv2D);
901 CHECK_SIZES(Conv2DBlockScaled);
902 CHECK_SIZES(Conv3D);
903 CHECK_SIZES(DepthwiseConv2D);
904 CHECK_SIZES(TransposeConv2D);
905 CHECK_SIZES(FFT2d);
906 CHECK_RANKS_AND_SIZES(MatMul);
907 CHECK_RANKS_AND_SIZES(MatMulT);
908 CHECK_SIZES(MatmulTBlockScaled);
909 CHECK_SIZES(MaxPool2d);
910 CHECK_SIZES(MaxPool2dAdaptive);
911 CHECK_SIZES(RFFT2d);
912 // Scatter/Gather Operators
914 CHECK_SIZES(RowGather);
915 CHECK_SIZES(Scatter);
916 // Image Operators
917 CHECK_SIZES(Resize);
918 // Custom Operators
919 CHECK_SIZES(Custom);
920 // Control Flow Operators
921 CHECK_SIZES(While);
922 // Shape Operators
923 CHECK_SIZES(ConstShape);
924
925 // For the following operations, check whether the shape length of each
926 // operand is valid given a level.
927
928 // Shape Operators
929 CHECK_SHAPE_LEN(AddShape);
930 CHECK_SHAPE_LEN(AssertEqualShape);
931 CHECK_SHAPE_LEN(ConcatShape);
932 CHECK_SHAPE_LEN(DivCeilShape);
933 CHECK_SHAPE_LEN(DivFloorShape);
934 CHECK_SHAPE_LEN(Exp2Shape);
935 CHECK_SHAPE_LEN(Log2CeilShape);
936 CHECK_SHAPE_LEN(Log2FloorShape);
937 CHECK_SHAPE_LEN(MaxShape);
938 CHECK_SHAPE_LEN(MinShape);
939 CHECK_SHAPE_LEN(ModShape);
940 CHECK_SHAPE_LEN(MulShape);
941 CHECK_SHAPE_LEN(SliceShape);
942 CHECK_SHAPE_LEN(SubShape);
943
944#undef CHECK_RANKS_AND_SIZES
945#undef CHECK_SIZES
946#undef CHECK_SHAPE_LEN
947 return success();
948}
949
950// Perform the Level tensor size check on the tensor type.
951LogicalResult TosaValidation::levelCheckSize(Operation *op,
952 const Type &typeToCheck,
953 const StringRef operandOrResult) {
954 if (ShapedType type = dyn_cast<ShapedType>(typeToCheck)) {
955 if (!type.hasRank())
956 return op->emitOpError() << "failed level check: unranked tensor";
957 auto shape = type.getShape();
958 for (auto dim : shape) {
959 const bool dimIsDynamic = mlir::ShapedType::isDynamic(dim);
960 const TosaSpecificationVersion targetVersion = targetEnv.getSpecVersion();
961 const TosaSpecificationVersion minRequiredVersion(1, 1, true);
962 if (targetVersion.isBackwardsCompatibleWith(minRequiredVersion) &&
963 dimIsDynamic)
964 // TOSA 1.1 and above supports dynamic dimensions, however, they must be
965 // resolved at backend compile time. Runtime dynamism is not currently
966 // supported. Checking this requirement is met is delegated to backends.
967 return success();
968
969 // When targeting TOSA 1.0 or below, dynamic dims are not supported
970 if (dimIsDynamic)
971 return op->emitOpError() << "failed level check: " << operandOrResult
972 << " shape dimension cannot be dynamic when"
973 << " targeting TOSA specification version 1.0"
974 << " or below";
975 }
976
977 int64_t elementBits = tosa::getBitWidth(getElementTypeOrSelf(type));
978 int64_t elementBytes = std::max(INT64_C(1), elementBits / 8);
979 int64_t size = elementBytes * type.getNumElements();
980
981 // According to 1.11. Tensor Definitions of Tosa spec, the value of
982 // tensor_size_t is 1 << MAX_LOG2_SIZE) - 1 where MAX_LOG2_SIZE is
983 // defined in 1.7. Levels.
984 // For each tensor, the number of tensor elements multiplied by the
985 // element size in bytes must be representable as a tensor_size_t.
986 const int64_t maxSize =
987 (INT64_C(1) << targetEnv.getLevel().MAX_LOG2_SIZE) - 1;
988 if (size > maxSize)
989 return op->emitOpError()
990 << "failed level check: " << operandOrResult
991 << " tensor size (in bytes) <= (1 << MAX_LOG2_SIZE - 1)";
992 }
993 return success();
994}
995
996LogicalResult TosaValidation::applyLevelCheck(Operation *op) {
997 if (targetEnv.getLevel() == TOSA_LEVEL_NONE) {
998 // no need to do level checks
999 return success();
1000 }
1001
1002 // check rank and sizes early so later checks can assume shaped operands
1003 if (failed(levelCheckRanksAndSizes(op)))
1004 return failure();
1005
1006 if (failed(levelCheckPool<tosa::AvgPool2dOp>(op)) ||
1007 failed(levelCheckAdaptivePool<tosa::AvgPool2dAdaptiveOp>(op)) ||
1008 failed(levelCheckConv<tosa::Conv2DOp>(op)) ||
1009 failed(levelCheckConv<tosa::Conv3DOp>(op)) ||
1010 failed(levelCheckConv<tosa::DepthwiseConv2DOp>(op)) ||
1011 failed(levelCheckFFT<tosa::FFT2dOp>(op)) ||
1012 failed(levelCheckPool<tosa::MaxPool2dOp>(op)) ||
1013 failed(levelCheckAdaptivePool<tosa::MaxPool2dAdaptiveOp>(op)) ||
1014 failed(levelCheckFFT<tosa::RFFT2dOp>(op)) ||
1015 failed(levelCheckTransposeConv2d(op)) || failed(levelCheckResize(op)) ||
1016 failed(levelCheckConv2DBlockScaled(op))) {
1017 return failure();
1018 }
1019
1020 // level check MAX_TENSOR_LIST_SIZE
1021 if (failed(levelCheckListSize(op))) {
1022 return failure();
1023 }
1024
1025 if (isa<tosa::IfOp>(op) || isa<tosa::WhileOp>(op)) {
1026 if (failed(levelCheckMaxNesting(op))) {
1027 return failure();
1028 }
1029 }
1030
1031 return success();
1032}
1033
1034LogicalResult TosaValidation::applyAttributeCheck(Operation *op) {
1035 if (failed(attributeCheckRescale(op)))
1036 return failure();
1037 if (failed(attributeCheckCast(op)))
1038 return failure();
1039 return success();
1040}
1041
1042inline bool CompatibleTypes(const mlir::Type &type,
1043 const mlir::Type &declaredType) {
1044 // for now, simply use type equality comparison
1045 return type == declaredType;
1046}
1047
1048LogicalResult TosaValidation::CheckVariable(Operation *op) {
1049 if (auto variableOp = dyn_cast<mlir::tosa::VariableOp>(op)) {
1050 mlir::StringAttr nameAttr = variableOp.getNameAttr();
1051
1052 if (variablesMap.count(nameAttr))
1053 return op->emitOpError() << "name has already been declared";
1054
1055 auto elementType = variableOp.getType();
1056 DenseIntElementsAttr varShapeAttr = variableOp.getVarShape();
1057 SmallVector<int64_t> shape = to_vector(varShapeAttr.getValues<int64_t>());
1058 RankedTensorType variableType =
1059 RankedTensorType::get(ArrayRef<int64_t>(shape), elementType);
1060
1061 variablesMap[nameAttr] = variableType;
1062 }
1063
1064 return success();
1065}
1066
1067LogicalResult TosaValidation::CheckVariableReadOrWrite(Operation *op) {
1068 if (isa<mlir::tosa::VariableReadOp>(op) ||
1069 isa<mlir::tosa::VariableWriteOp>(op)) {
1070 mlir::StringAttr nameAttr =
1072 .Case<mlir::tosa::VariableReadOp, mlir::tosa::VariableWriteOp>(
1073 [](auto variableOp) { return variableOp.getNameAttr(); });
1074 if (!variablesMap.count(nameAttr))
1075 return op->emitOpError() << "name has not been declared";
1076
1077 auto varType = variablesMap[nameAttr];
1078
1079 for (auto v : op->getOperands()) {
1080 auto type = v.getType();
1081 if (!CompatibleTypes(type, varType))
1082 return op->emitOpError() << "operand type does not equal variable type";
1083 }
1084
1085 for (auto v : op->getResults()) {
1086 auto type = v.getType();
1087 if (!CompatibleTypes(type, varType))
1088 return op->emitOpError() << "result type does not equal variable type";
1089 }
1090 }
1091
1092 return success();
1093}
1094
1095LogicalResult TosaValidation::applyVariableCheck(Operation *op) {
1096 if (failed(CheckVariable(op)) || failed(CheckVariableReadOrWrite(op)))
1097 return failure();
1098 return success();
1099}
1100
1101LogicalResult checkErrorIfResize(Operation *op) {
1102 auto resize = dyn_cast<tosa::ResizeOp>(op);
1103 if (!resize)
1104 return success();
1105
1106 const Value input = resize.getInput();
1107 const Value output = resize.getOutput();
1108 const RankedTensorType inputType =
1109 llvm::dyn_cast<RankedTensorType>(input.getType());
1110 const RankedTensorType outputType =
1111 llvm::dyn_cast<RankedTensorType>(output.getType());
1112
1113 if (!inputType || !outputType)
1114 return op->emitOpError("expect ranked input/output tensor");
1115
1116 // Ensure the image size is supported by GPU APIs and that for integer
1117 // implementations, position * stride does not overflow int32_t.
1118 if (inputType.hasStaticShape() && outputType.hasStaticShape()) {
1119 const SmallVector<int64_t, 4> sizes = {
1120 outputType.getDimSize(1), outputType.getDimSize(2),
1121 inputType.getDimSize(1), inputType.getDimSize(2)};
1122 const int64_t *maxDim = llvm::max_element(sizes);
1123 if (maxDim != sizes.end() && *maxDim >= 16384)
1124 return op->emitOpError(
1125 "expect input/output height/width dims to be < 16384, ")
1126 << "got [OH, OW, IH, IW] = " << sizes;
1127 }
1128
1129 SmallVector<int64_t> scale;
1130 if (!tosa::getConstShapeValues(resize.getScale().getDefiningOp(), scale))
1131 return failure();
1132
1133 const int64_t scaleYN = scale[0];
1134 const int64_t scaleYD = scale[1];
1135 const int64_t scaleXN = scale[2];
1136 const int64_t scaleXD = scale[3];
1137
1138 // Ensure scale values don't overflow int32 accumulator
1139 if (scaleYN > (1 << 11) || scaleXN > (1 << 11))
1140 return op->emitOpError(
1141 "expect all scale numerator values to be <= (1 << 11), "
1142 "got scale_y_n=")
1143 << scaleYN << ", scale_x_n=" << scaleXN;
1144
1145 if (scaleYD >= 16 * scaleYN || scaleXD >= 16 * scaleXN)
1146 return op->emitOpError("expect a downscale ratio larger than 1/16, got y=")
1147 << scaleYN << "/" << scaleYD << ", x=" << scaleXN << "/" << scaleXD;
1148
1149 SmallVector<int64_t> offset;
1150 SmallVector<int64_t> border;
1151 if (!tosa::getConstShapeValues(resize.getOffset().getDefiningOp(), offset) ||
1152 !tosa::getConstShapeValues(resize.getBorder().getDefiningOp(), border))
1153 return failure();
1154
1155 const int64_t offsetY = offset[0];
1156 const int64_t offsetX = offset[1];
1157 // Set a consistent lower limit of 1/16 downscale to simplify
1158 // implementations
1159 if (offsetY < -scaleYN || offsetY >= 16 * scaleYN)
1160 return op->emitOpError(
1161 "expect offsetY / scaleYNumerator to be in range [-1, 16), got ")
1162 << offsetY << "/" << scaleYN;
1163 if (offsetX < -scaleXN || offsetX >= 16 * scaleXN)
1164 return op->emitOpError(
1165 "expect offsetX / scaleXNumerator to be in range [-1, 16), got ")
1166 << offsetX << "/" << scaleXN;
1167
1168 const int64_t borderY = border[0];
1169 const int64_t borderX = border[1];
1170 if (borderY < -16 * scaleYN || borderY >= scaleYN)
1171 return op->emitOpError(
1172 "expect borderY / scaleYNumerator to be in range [-16, 1), got ")
1173 << borderY << "/" << scaleYN;
1174 if (borderX < -16 * scaleXN || borderX >= scaleXN)
1175 return op->emitOpError(
1176 "expect borderX / scaleXNumerator to be in range [-16, 1), got ")
1177 << borderX << "/" << scaleXN;
1178
1179 // The following section of code is mostly duplicated with ResizeOp::verify().
1180 //
1181 // In TOSA specification, we do not support broadcast behavior.
1182 // However, there is a rewrite pattern to materialize broadcast ResizeOp.
1183 // It makes invalid TOSA ResizeOp into valid one. To avoid breaking
1184 // existing code, we keep the rewrite pattern untouched. So, we need
1185 // loose the checking in ResizeOp::verify() to support broadcast ResizeOp.
1186 //
1187 // Here is a strict checking to conform TOSA specification.
1188 // FIXME: Remove the duplicated checkings when broadcast ResizeOp is removed.
1189 auto idivCheck = [](const int64_t lhs,
1190 const int64_t rhs) -> std::optional<int64_t> {
1191 if (lhs % rhs != 0)
1192 return std::nullopt;
1193 return lhs / rhs;
1194 };
1195
1196 const int64_t oh = outputType.getDimSize(1);
1197 const int64_t ow = outputType.getDimSize(2);
1198 const int64_t ih = inputType.getDimSize(1);
1199 const int64_t iw = inputType.getDimSize(2);
1200
1201 if (ih != ShapedType::kDynamic) {
1202 const std::optional<int64_t> calculatedOutHeightMinusOne =
1203 idivCheck((ih - 1) * scaleYN - offsetY + borderY, scaleYD);
1204 if (!calculatedOutHeightMinusOne.has_value())
1205 return op->emitOpError(
1206 "expected (input_height - 1) * scale_y_n - offset_y + "
1207 "border_y ")
1208 << "to be wholly divisible by scale_y_d, got ((" << ih
1209 << " - 1) * " << scaleYN << " - " << offsetY << " + " << borderY
1210 << ") / " << scaleYD;
1211 const int64_t calculatedOutHeight = calculatedOutHeightMinusOne.value() + 1;
1212 if (oh != ShapedType::kDynamic && calculatedOutHeight != oh)
1213 return op->emitOpError(
1214 "calculated output height did not match expected: ")
1215 << "calculated=" << calculatedOutHeight << ", expected=" << oh;
1216 }
1217
1218 if (iw != ShapedType::kDynamic) {
1219 const std::optional<int64_t> calculatedOutWidthMinusOne =
1220 idivCheck((iw - 1) * scaleXN - offsetX + borderX, scaleXD);
1221 if (!calculatedOutWidthMinusOne.has_value())
1222 return op->emitOpError(
1223 "expected (input_width - 1) * scale_x_n - offset_x + "
1224 "border_x ")
1225 << "to be wholly divisible by scale_x_d, got ((" << iw
1226 << " - 1) * " << scaleXN << " - " << offsetX << " + " << borderX
1227 << ") / " << scaleXD;
1228 const int64_t calculatedOutWidth = calculatedOutWidthMinusOne.value() + 1;
1229 if (ow != ShapedType::kDynamic && calculatedOutWidth != ow)
1230 return op->emitOpError("calculated output width did not match expected: ")
1231 << "calculated=" << calculatedOutWidth << ", expected=" << ow;
1232 }
1233
1234 return success();
1235}
1236
1237LogicalResult checkErrorIfMul(Operation *op) {
1238 auto mul = dyn_cast<tosa::MulOp>(op);
1239 if (!mul)
1240 return success();
1241
1242 // REQUIRE(0 <= shift && shift <= 63);
1243 // REQUIRE(is_same<in_t,int32_t>() || shift == 0);
1244 ElementsAttr shift_elem;
1245 if (!matchPattern(mul.getShift(), m_Constant(&shift_elem)))
1246 return success();
1247 int32_t shift = shift_elem.getValues<IntegerAttr>()[0].getInt();
1248 auto inputElemType = getElementTypeOrSelf(mul.getInput1());
1249 if (inputElemType.isInteger(32)) {
1250 // 0 <= shift <= 63 for int32_t type
1251 if (shift < 0 || shift > 63)
1252 return op->emitOpError()
1253 << "requires 0 <= shift && shift <= 63, but got: " << shift;
1254 } else {
1255 // shift must be 0 for all other types
1256 if (shift != 0)
1257 return op->emitOpError()
1258 << "requires shift = 0 for all input data types that "
1259 "are not int32_t, but got: "
1260 << shift;
1262
1263 return success();
1265
1266LogicalResult checkErrorIfTable(Operation *op) {
1267 auto table = dyn_cast<tosa::TableOp>(op);
1268 if (!table)
1269 return success();
1270
1271 // REQUIRE(length(table) == TABLE_SIZE) where TABLE_SIZE is 256 or 513
1272 const auto inputElemType = getElementTypeOrSelf(table.getInput1().getType());
1273 const int tableSize = inputElemType.isInteger(8) ? 256 : 513;
1274
1275 const ShapeAdaptor tableShape(table.getTable().getType());
1276 if (tableShape.hasStaticShape()) {
1277 const auto numElements = tableShape.getNumElements();
1278 if (numElements != tableSize)
1279 return op->emitOpError() << "requires table size of " << tableSize
1280 << ", got " << numElements;
1281 }
1282
1283 return success();
1284}
1286LogicalResult checkErrorIfRescale(Operation *op) {
1287 auto rescale = dyn_cast<tosa::RescaleOp>(op);
1288 if (!rescale)
1289 return success();
1290
1291 auto inputType = llvm::dyn_cast<ShapedType>(rescale.getInput().getType());
1292 auto outputType = llvm::dyn_cast<ShapedType>(rescale.getOutput().getType());
1293 if (!inputType || !outputType || !inputType.getElementType().isInteger() ||
1294 !outputType.getElementType().isInteger())
1295 return success();
1296
1297 auto inElemType = inputType.getElementType();
1298 auto outElemType = outputType.getElementType();
1299 auto inWidth = inElemType.getIntOrFloatBitWidth();
1300 auto outWidth = outElemType.getIntOrFloatBitWidth();
1302 bool inputUnsigned = rescale.getInputUnsigned();
1303 bool outputUnsigned = rescale.getOutputUnsigned();
1304
1305 bool scale32 = rescale.getScale32();
1306 auto roundingMode = rescale.getRoundingMode();
1307
1308 // ERROR_IF(scale32 && is_same<in_t,i48_t>())
1309 if (scale32 && inWidth == 48)
1310 return op->emitOpError() << "scale32 is not allowed with 48-bit input.";
1311
1312 // ERROR_IF(!scale32 && (rounding_mode == DOUBLE_ROUND))
1313 if (!scale32 && roundingMode == RoundingMode::DOUBLE_ROUND)
1314 return op->emitOpError()
1315 << "DOUBLE_ROUND is only allowed with scale32=true.";
1316
1317 // ERROR_IF(input_unsigned && output_unsigned)
1318 if (inputUnsigned && outputUnsigned)
1319 return op->emitOpError() << "input and output cannot be both unsigned.";
1320
1321 // ERROR_IF(is_same<out_t,i32_t>() && input_unsigned)
1322 if (outWidth == 32 && inputUnsigned)
1323 return op->emitOpError()
1324 << "i32 output type is not allowed with unsigned input.";
1325
1326 // ERROR_IF(is_same<in_t,i32_t>() && output_unsigned)
1327 if (inWidth == 32 && outputUnsigned)
1328 return op->emitOpError()
1329 << "i32 input type is not allowed with unsigned output.";
1330
1331 // ERROR_IF(is_same<in_t,i48_t>() && output_unsigned)
1332 if (inWidth == 48 && outputUnsigned)
1333 return op->emitOpError()
1334 << "i48 input type is not allowed with unsigned output.";
1335
1336 // ERROR_IF(is_same<in_t, i48_t> && input_unsigned)
1337 if (inWidth == 48 && inputUnsigned)
1338 return op->emitOpError() << "i48 input type cannot be unsigned.";
1339
1340 // ERROR_IF(is_same<in_t, i32_t> && input_unsigned)
1341 if (inWidth == 32 && inputUnsigned)
1342 return op->emitOpError() << "i32 input type cannot be unsigned.";
1343
1344 // ERROR_IF(is_same<out_t, i32_t> && output_unsigned)
1345 if (outWidth == 32 && outputUnsigned)
1346 return op->emitOpError() << "i32 output type cannot be unsigned.";
1347
1348 return success();
1349}
1350
1351LogicalResult checkErrorIfPad(Operation *op) {
1352 auto pad = dyn_cast<tosa::PadOp>(op);
1353 if (!pad)
1354 return success();
1355
1356 DenseIntElementsAttr paddingAttr;
1357 if (!matchPattern(pad.getPadding(), m_Constant(&paddingAttr)))
1358 // Pad verifier will catch this
1359 return success();
1360
1361 for (const APInt &val : paddingAttr.getValues<APInt>()) {
1362 if (val.getSExtValue() < 0)
1363 return op->emitOpError() << "padding value must all be non-negative, got "
1364 << val.getSExtValue();
1365 }
1366
1367 return success();
1368}
1369
1370LogicalResult checkErrorIfReshape(Operation *op) {
1371 auto reshapeOp = dyn_cast<tosa::ReshapeOp>(op);
1372 if (!reshapeOp)
1373 return success();
1374
1375 SmallVector<int64_t> shapeValues;
1376 if (!tosa::getConstShapeValues(reshapeOp.getShape().getDefiningOp(),
1377 shapeValues))
1378 return success();
1379
1380 if (llvm::is_contained(shapeValues, kInferableDimSize))
1381 return op->emitOpError("shape input contains inferable dimension (")
1383 << ") "
1384 "which does not conform to the TOSA specification";
1385
1386 return success();
1387}
1388
1389LogicalResult checkErrorIfSlice(Operation *op) {
1390 auto sliceOp = dyn_cast<tosa::SliceOp>(op);
1391 if (!sliceOp)
1392 return success();
1393
1394 SmallVector<int64_t> startValues;
1395 SmallVector<int64_t> sizeValues;
1396 const bool hasStartValues = tosa::getConstShapeValues(
1397 sliceOp.getStart().getDefiningOp(), startValues);
1398 const bool hasSizeValues =
1399 tosa::getConstShapeValues(sliceOp.getSize().getDefiningOp(), sizeValues);
1400
1401 if (hasStartValues && llvm::is_contained(startValues, kInferableDimSize))
1402 return op->emitOpError("start input contains inferable dimension (")
1404 << ") which does not conform to the TOSA specification";
1405 if (hasSizeValues && llvm::is_contained(sizeValues, kInferableDimSize))
1406 return op->emitOpError("size input contains inferable dimension (")
1408 << ") which "
1409 "does not conform to the TOSA specification";
1410
1411 return success();
1412}
1413
1414static bool isOpIsolatedWithinRegion(Operation *op, Region *region) {
1415 return llvm::all_of(op->getOperands(), [&](auto operand) {
1416 Region *operandRegion = operand.getParentRegion();
1417 return operandRegion && region->isAncestor(operandRegion);
1418 });
1419}
1420
1421static LogicalResult isRegionIsolatedFromAbove(Region &regionToCheck) {
1422 bool noLiveInValue = true;
1423 regionToCheck.walk([&noLiveInValue, &regionToCheck](Operation *op) {
1424 if (!isOpIsolatedWithinRegion(op, &regionToCheck)) {
1425 noLiveInValue = false;
1426 return WalkResult::interrupt();
1427 }
1428 return WalkResult::advance();
1429 });
1430 return noLiveInValue ? success() : failure();
1431}
1432
1433LogicalResult checkIsolatedRegion(Operation *op, Region &regionToCheck,
1434 StringRef regionName) {
1435 if (succeeded(isRegionIsolatedFromAbove(regionToCheck)))
1436 return success();
1437 return op->emitOpError()
1438 << "is not conformant to the TOSA specification. It requires the '"
1439 << regionName << "' region is isolated from above.\n";
1440}
1441
1442LogicalResult checkErrorIfCondIf(Operation *op) {
1443 auto ifOp = dyn_cast<tosa::IfOp>(op);
1444 if (!ifOp)
1445 return success();
1446
1447 // Currently the dialect supports declaring cond_if operations that
1448 // have then/else regions that reference values from outside these
1449 // regions. According to the specification, all values used by the
1450 // then/else regions must be explicitly declared within the regions.
1451 // Therefore we must check that the then/else regions are
1452 // "isolated from above", in order to be conformant to the
1453 // specification.
1454 //
1455 // Note: the dialect currently supports two styles of syntax for
1456 // declaring "cond_if" operations. We'll refer to these as follows:
1457 //
1458 // Generic:
1459 // %0 = "tosa.cond_if"(%arg0, %arg1, %arg2) ({
1460 // ^bb0(%arg3, %arg4):
1461 // tosa.yield %arg3
1462 // }, {
1463 // ^bb0(%arg3, %arg4):
1464 // tosa.yield %arg4
1465 // })
1466 //
1467 // Simplified:
1468 // %0 = tosa.cond_if %arg2 (%arg3 = %arg0, %arg4 = %arg1) {
1469 // ^bb0(%arg3, %arg4):
1470 // tosa.yield %arg3
1471 // } else {
1472 // ^bb0(%arg3, %arg4):
1473 // tosa.yield %arg4
1474 // }
1475
1476 if (failed(checkIsolatedRegion(op, ifOp.getThenGraph(), "then")) ||
1477 failed(checkIsolatedRegion(op, ifOp.getElseGraph(), "else")))
1478 return failure();
1479 return success();
1480}
1481
1482LogicalResult checkErrorIfWhileLoop(Operation *op) {
1483 auto whileOp = dyn_cast<tosa::WhileOp>(op);
1484 if (!whileOp)
1485 return success();
1486
1487 if (failed(checkIsolatedRegion(op, whileOp.getCondGraph(), "cond")) ||
1488 failed(checkIsolatedRegion(op, whileOp.getBodyGraph(), "body")))
1489 return failure();
1490 return success();
1491}
1492
1493LogicalResult checkErrorIfScatter(Operation *op) {
1494 auto scatterOp = dyn_cast<tosa::ScatterOp>(op);
1495 if (!scatterOp)
1496 return success();
1497
1498 // for constant indices, check that there are no duplicate values
1499 DenseIntElementsAttr indicesAttr;
1500 if (!matchPattern(scatterOp.getIndices(), m_Constant(&indicesAttr)))
1501 return success();
1502
1503 auto const indicesType =
1504 dyn_cast<ShapedType>(scatterOp.getIndices().getType());
1505 if (!indicesType || !indicesType.hasRank()) {
1506 op->emitOpError("expect ranked indices tensor");
1507 return failure();
1508 }
1509
1510 if (!hasUniqueConstantScatterIndices(indicesType, indicesAttr)) {
1511 op->emitOpError("indices values contain duplicates");
1512 return failure();
1513 }
1514
1515 return success();
1516}
1517
1518LogicalResult TosaValidation::applyErrorIfCheck(Operation *op) {
1519 if (failed(checkErrorIfResize(op)) || failed(checkErrorIfMul(op)) ||
1520 failed(checkErrorIfTable(op)) || failed(checkErrorIfRescale(op)) ||
1521 failed(checkErrorIfPad(op)) || failed(checkErrorIfReshape(op)) ||
1522 failed(checkErrorIfSlice(op)) || failed(checkErrorIfCondIf(op)) ||
1523 failed(checkErrorIfWhileLoop(op)) || failed(checkErrorIfScatter(op)))
1524 return failure();
1525 return success();
1526}
1527
1528LogicalResult TosaValidation::applyFunctionSignatureCheck(func::FuncOp op) {
1529 // Require tensor type parameters and results
1530 const auto isTensorType = [](Type type) { return isa<TensorType>(type); };
1531 if (!llvm::all_of(op.getArgumentTypes(), isTensorType))
1532 return op.emitOpError()
1533 << "Function argument types must be a tensor type to be TOSA "
1534 "compliant, got !tosa.shape type";
1535 if (!llvm::all_of(op.getResultTypes(), isTensorType))
1536 return op.emitOpError()
1537 << "Function return types must be a tensor type to be TOSA "
1538 "compliant, got !tosa.shape type";
1539
1540 // Validate element types
1541 if (failed(validateOperationElementTypes(op, !strictOpSpecAlignment)))
1542 return failure();
1543
1544 // Level check
1545 const TosaLevel tosaLevel = targetEnv.getLevel();
1546 for (const auto &[idx, argType] : llvm::enumerate(op.getArgumentTypes())) {
1547 const std::string inputDesc = llvm::formatv("input argument {0}", idx);
1548 if (failed(levelCheckRank(op, argType, inputDesc, tosaLevel.MAX_RANK)))
1549 return failure();
1550 if (failed(levelCheckSize(op, argType, inputDesc)))
1551 return failure();
1552 }
1553 for (const auto &[idx, resultType] : llvm::enumerate(op.getResultTypes())) {
1554 const std::string resultDesc = llvm::formatv("return value {0}", idx);
1555 if (failed(levelCheckRank(op, resultType, resultDesc, tosaLevel.MAX_RANK)))
1556 return failure();
1557 if (failed(levelCheckSize(op, resultType, resultDesc)))
1558 return failure();
1559 }
1560
1561 // Explicitly check for no zero dimensions
1562 // Note: This check is not required for TOSA operations since it is mandated
1563 // on construction
1564 for (const Type &argType :
1565 llvm::concat<const Type>(op.getArgumentTypes(), op.getResultTypes())) {
1566 if (auto shapedType = dyn_cast<ShapedType>(argType)) {
1567 if (llvm::any_of(shapedType.getShape(),
1568 [](int64_t dim) { return dim == 0; }))
1569 return op.emitOpError() << "Function argument or return types must not "
1570 "have zero dimensions";
1571 }
1572 }
1573
1574 return success();
1575}
1576
1577LogicalResult TosaValidation::validateValidElementType(Operation *op, Type type,
1578 bool allowUnsigned) {
1579 if (isa<FloatType>(type)) {
1580 if (isa<Float32Type, Float16Type, BFloat16Type, Float8E4M3FNType,
1581 Float8E5M2Type, Float4E2M1FNType, Float6E2M3FNType,
1582 Float6E3M2FNType, Float8E8M0FNUType>(type))
1583 return success();
1584 } else if (auto intTy = dyn_cast<IntegerType>(type)) {
1585 if (intTy.isSignless()) {
1586 switch (intTy.getWidth()) {
1587 case 1:
1588 case 4:
1589 case 8:
1590 case 16:
1591 case 32:
1592 case 48:
1593 case 64:
1594 return success();
1595 }
1596 } else if (allowUnsigned && intTy.isUnsigned()) {
1597 switch (intTy.getWidth()) {
1598 case 8:
1599 case 16:
1600 case 32:
1601 return success();
1602 }
1603 }
1604 } else if (isa<tosa::shapeType>(type))
1605 return success();
1606 else if (isa<tosa::mxint8Type, tosa::BlockScaledType>(type))
1607 return success();
1608
1609 return op->emitOpError() << "is not profile-aligned: element type " << type
1610 << " is not legal";
1611}
1612
1613LogicalResult
1614TosaValidation::validateOperationElementTypes(TosaOp op, bool allowUnsigned) {
1615 for (Value operand : op->getOperands()) {
1616 Type elementTy = getElementTypeOrSelf(operand);
1617 if (failed(validateValidElementType(op, elementTy, allowUnsigned)))
1618 return failure();
1619 }
1620
1621 for (Type resultTy : op->getResultTypes()) {
1622 Type elementTy = getElementTypeOrSelf(resultTy);
1623 if (failed(validateValidElementType(op, elementTy, allowUnsigned)))
1624 return failure();
1625 }
1626
1627 if (auto variableOp = dyn_cast<tosa::VariableOp>(*op)) {
1628 if (failed(
1629 validateValidElementType(op, variableOp.getType(), allowUnsigned)))
1630 return failure();
1631 }
1632 return success();
1633}
1634
1635LogicalResult
1636TosaValidation::validateOperationElementTypes(func::FuncOp op,
1637 bool allowUnsigned) {
1638 for (const Type &argType :
1639 llvm::concat<const Type>(op.getArgumentTypes(), op.getResultTypes())) {
1640 const Type elementTy = getElementTypeOrSelf(argType);
1641 if (failed(validateValidElementType(op, elementTy, allowUnsigned)))
1642 return failure();
1643 }
1644
1645 return success();
1646}
1647
1648void TosaValidation::runOnOperation() {
1649 ModuleOp modOp = getOperation();
1650 TosaDialect *tosaDialect = getContext().getLoadedDialect<TosaDialect>();
1651 if (!tosaDialect)
1652 return;
1653
1654 const TargetEnvAttr targetEnvAttr = lookupTargetEnvOrDefault(modOp);
1655 const auto maybeTargetEnv =
1656 tosa::TargetEnv::createTargetEnvFromAttr(targetEnvAttr, modOp.getLoc());
1657 if (failed(maybeTargetEnv))
1658 return signalPassFailure();
1659 targetEnv = *maybeTargetEnv;
1660
1661 const auto functions = modOp.getOps<func::FuncOp>();
1662 if (validateFunctionSignature &&
1663 llvm::any_of(functions, [&](func::FuncOp func) {
1664 return failed(applyFunctionSignatureCheck(func));
1665 }))
1666 return signalPassFailure();
1667
1668 modOp.walk([&](TosaOp op) {
1669 // validate operator element types:
1670 // - rescale operator is allowed to have ui8/ui16/ui32
1671 // operands/results when strictOpSpecAlignment is false
1672 // - perform valid element type check at the beginning to
1673 // protect rest of code against quantized element types
1674 const bool allowUnsigned =
1675 !strictOpSpecAlignment && isa<tosa::RescaleOp>(op);
1676 if (failed(validateOperationElementTypes(op, allowUnsigned)))
1677 return signalPassFailure();
1678
1679 if (strictOpSpecAlignment &&
1680 failed(profileComp.checkProfile(op, targetEnv)))
1681 return signalPassFailure();
1682
1683 if (strictOpSpecAlignment &&
1684 failed(profileComp.checkExtension(op, targetEnv)))
1685 return signalPassFailure();
1686
1687 if (strictOpSpecAlignment &&
1688 failed(checkSpecificationVersionConstraint(op, targetEnv)))
1689 return signalPassFailure();
1690
1691 if (!allowInvalidOpDatatypeCombinations &&
1692 failed(profileComp.checkInvalid(op)))
1693 return signalPassFailure();
1694
1695 // Some uses of TOSA rely on the constant operands of particular
1696 // operations.
1697 if (failed(applyConstantOperandCheck(op)))
1698 signalPassFailure();
1699
1700 // do level checks
1701 if (failed(applyLevelCheck(op)))
1702 signalPassFailure();
1703
1704 // check additional attribute restrictions
1705 if (failed(applyAttributeCheck(op)))
1706 signalPassFailure();
1707
1708 // do variable type checks
1709 if (failed(applyVariableCheck(op)))
1710 signalPassFailure();
1711
1712 // do error if checks
1713 if (strictOpSpecAlignment && failed(applyErrorIfCheck(op)))
1714 signalPassFailure();
1715 });
1716}
1717} // namespace
return success()
lhs
b getContext())
static llvm::ManagedStatic< PassManagerOptions > options
static std::optional< int64_t > idivCheck(const int64_t lhs, const int64_t rhs)
Definition TosaOps.cpp:279
#define CHECK_RANKS_AND_SIZES(tosaOp)
#define CHECK_SIZES(tosaOp)
#define CHECK_SHAPE_LEN(tosaOp)
@ Gather
#define mul(a, b)
LogicalResult checkProfile(Operation *op, const tosa::TargetEnv &targetEnv)
LogicalResult checkExtension(Operation *op, const tosa::TargetEnv &targetEnv)
LogicalResult checkInvalid(Operation *op)
Attributes are known-constant values of operations.
Definition Attributes.h:25
An attribute that represents a reference to a dense integer vector or tensor object.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
result_range getResults()
Definition Operation.h:440
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
RetT walk(FnT &&callback)
Walk all nested operations, blocks or regions (including this region), depending on the type of callb...
Definition Region.h:297
Adaptor class to abstract the differences between whether value is from a ShapedType or ShapedTypeCom...
Type getType() const
Return the type of this value.
Definition Value.h:105
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
This class represents the capability enabled in the target implementation such as profile,...
Definition TargetEnv.h:119
TosaLevel getLevel() const
Definition TargetEnv.h:136
static FailureOr< TargetEnv > createTargetEnvFromAttr(TargetEnvAttr targetAttr, Location targetEnvAttrLoc)
bool allows(Profile prof) const
Definition TargetEnv.h:139
TosaSpecificationVersion getSpecVersion() const
Definition TargetEnv.h:132
A thin wrapper around the SpecificationVersion enum to represent and provide utilities around the TOS...
Definition TargetEnv.h:60
bool isBackwardsCompatibleWith(TosaSpecificationVersion baseVersion) const
Definition TargetEnv.h:69
SmallVector< AffineExpr, 4 > concat(ArrayRef< AffineExpr > a, ArrayRef< AffineExpr > b)
Return the vector that is the concatenation of a and b.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
llvm::SmallString< 4 > stringifyVersion(TosaSpecificationVersion version)
Definition TargetEnv.cpp:25
RankedTensorType getVariableType(VariableOp variableOp)
static constexpr TosaLevel TOSA_LEVEL_NONE
Definition TargetEnv.h:45
bool hasUniqueConstantScatterIndices(ShapedType indicesType, DenseIntElementsAttr indicesAttr)
constexpr int64_t kInferableDimSize
Represents a dimension in the shape of a tensor that can be inferred based on the other provided dime...
Definition TosaOps.h:102
unsigned getBitWidth(Type type)
Definition TosaOps.cpp:331
TargetEnvAttr lookupTargetEnvOrDefault(Operation *op)
Queries the target environment recursively from enclosing symbol table ops containing the given op or...
bool getConstShapeValues(Operation *op, llvm::SmallVector< int64_t > &result_shape)
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
@ Mul
RHS of mul is always a constant or a symbolic expression.
Definition AffineExpr.h:43
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
LogicalResult verifyCompatibleDims(ArrayRef< int64_t > dims)
Dimensions are compatible if all non-dynamic dims are equal.
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369