30#include "llvm/ADT/STLExtras.h"
31#include "llvm/ADT/StringExtras.h"
32#include "llvm/ADT/TypeSwitch.h"
33#include "llvm/Support/FormatVariadic.h"
37#define GEN_PASS_DEF_TOSAVALIDATION
38#include "mlir/Dialect/Tosa/Transforms/Passes.h.inc"
49 for (
const auto index : operandIndices) {
52 return op->
emitOpError(
"expected compile time resolvable constant, but "
53 "got variable value for operand #")
60static LogicalResult checkConstantOperandMul(
Operation *op,
62 if (!env.
allows(Extension::dynamic) && isa<tosa::MulOp>(op)) {
64 return checkConstantOperands(op, {2});
69static LogicalResult checkConstantOperandTable(
Operation *op,
71 if (!env.
allows(Extension::dynamic) && isa<tosa::TableOp>(op)) {
73 return checkConstantOperands(op, {1});
78static LogicalResult checkConstantOperandPad(
Operation *op,
80 if (
auto padOp = dyn_cast<tosa::PadOp>(op)) {
82 if (!env.
allows(Extension::dynamic) && padOp.getPadConst())
85 return checkConstantOperands(op, {2});
90static LogicalResult checkConstantOperandRescale(
Operation *op,
92 if (!env.
allows(Extension::dynamic) && isa<tosa::RescaleOp>(op)) {
94 return checkConstantOperands(op, {1, 2, 3, 4});
100static LogicalResult checkConstantOperandConvOps(
Operation *op,
102 if (!env.
allows(Extension::dynamic) && isa<T>(op)) {
104 return checkConstantOperands(op, {3, 4});
109static LogicalResult checkConstantOperandMatMul(
Operation *op,
111 if (!env.
allows(Extension::dynamic) &&
112 isa<tosa::MatMulOp, tosa::MatMulTOp>(op)) {
114 return checkConstantOperands(op, {2, 3});
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});
130static LogicalResult checkConstantOperandRowGather(
Operation *op,
132 if (!env.
allows(Extension::dynamic) && isa<tosa::RowGatherOp>(op)) {
134 return checkConstantOperands(op, {2});
139static LogicalResult checkConstantOperandAvgPool2d(
Operation *op,
141 if (!env.
allows(Extension::dynamic) && isa<tosa::AvgPool2dOp>(op)) {
143 return checkConstantOperands(op, {1, 2});
150 if (!env.
allows(Extension::dynamic) && isa<tosa::AvgPool2dAdaptiveOp>(op)) {
154 return checkConstantOperands(op, {1, 2});
159static LogicalResult checkConstantOperandNegate(
Operation *op,
161 if (!env.
allows(Extension::dynamic) && isa<tosa::NegateOp>(op)) {
163 return checkConstantOperands(op, {1, 2});
168static LogicalResult checkConstantOperandSilceShape(
Operation *op,
170 if (!env.
allows(Extension::dynamic) && isa<tosa::SliceShapeOp>(op)) {
172 return checkConstantOperands(op, {1, 2});
179static LogicalResult checkSpecificationVersionConstraint(
Operation *op,
181 auto matmul = dyn_cast<tosa::MatMulOp>(op);
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;
193 if ((!aType || aType.getRank() == 3) && (!bType || bType.getRank() == 3) &&
194 (!outputType || outputType.getRank() == 3) &&
196 getBatchDimOrDynamic(bType),
197 getBatchDimOrDynamic(outputType)})))
201 "MATMUL ranks other than 3 or batch broadcasting require TOSA "
202 "specification version 1.1.draft");
211 explicit TosaValidation() { populateConstantOperandChecks(); }
213 explicit TosaValidation(
const TosaValidationOptions &
options)
215 this->strictOpSpecAlignment =
options.strictOpSpecAlignment;
216 this->allowInvalidOpDatatypeCombinations =
217 options.allowInvalidOpDatatypeCombinations;
218 this->validateFunctionSignature =
options.validateFunctionSignature;
220 void runOnOperation() final;
222 LogicalResult applyConstantOperandCheck(Operation *op) {
223 for (
auto &checker : constCheckers) {
224 if (
failed(checker(op, targetEnv)))
230 LogicalResult applyFunctionSignatureCheck(func::FuncOp op);
231 LogicalResult applyLevelCheck(Operation *op);
232 LogicalResult applyAttributeCheck(Operation *op);
235 LogicalResult applyVariableCheck(Operation *op);
238 LogicalResult applyErrorIfCheck(Operation *op);
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);
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)
266 <<
"failed level check: " << inputName <<
" <= " << levelName
267 <<
" (" << maxLevel <<
"), got " << calculatedValue;
271 LogicalResult levelCheckKernel(Operation *op, int32_t v,
272 const StringRef inputName) {
273 return levelCheck(op, v, targetEnv.getLevel().MAX_KERNEL, inputName,
277 LogicalResult levelCheckStride(Operation *op, int32_t v,
278 const StringRef inputName) {
279 return levelCheck(op, v, targetEnv.getLevel().MAX_STRIDE, inputName,
283 LogicalResult levelCheckScale(Operation *op, int32_t v,
284 const StringRef inputName) {
285 return levelCheck(op, v, targetEnv.getLevel().MAX_SCALE, inputName,
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");
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)) {
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";
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);
319 LogicalResult levelCheckSize(Operation *op,
const Type &typeToCheck,
320 const StringRef operandOrResult);
323 LogicalResult levelCheckSize(Operation *op,
const Value &v,
324 const StringRef operandOrResult) {
325 return levelCheckSize(op, v.
getType(), operandOrResult);
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)
334 <<
"failed shape type level check: " << typeToCheck
335 <<
" exceeds MAX_SHAPE_LEN";
341 template <
typename T>
342 LogicalResult levelCheckSizes(T tosaOp) {
343 auto op = tosaOp.getOperation();
345 if (
failed(levelCheckSize(op, v,
"operand")))
350 if (
failed(levelCheckSize(op, v,
"result")))
357 template <
typename T>
358 LogicalResult levelCheckRanks(T tosaOp) {
359 auto op = tosaOp.getOperation();
360 const TosaLevel tosaLevel = targetEnv.getLevel();
374 template <
typename T>
375 LogicalResult levelCheckShapeLengths(T tosaOp) {
376 for (
const auto &v : tosaOp->getOperands()) {
377 if (
failed(levelCheckShapeLength(tosaOp, v.getType(),
"operand")))
380 for (
const auto &v : tosaOp->getResults()) {
381 if (
failed(levelCheckShapeLength(tosaOp, v.getType(),
"result")))
389 LogicalResult levelCheckRanksAndSizes(Operation *op);
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"))) {
400 for (
auto s : poolOp.getStride()) {
401 if (
failed(levelCheckStride(op, s,
"stride"))) {
405 for (
auto p : poolOp.getPad()) {
406 if (
failed(levelCheckKernel(op, p,
"pad"))) {
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>;
419 template <
typename T,
typename std::enable_if<IsSupportedAdaptivePoolOp<T>,
421 LogicalResult levelCheckAdaptivePool(Operation *op) {
422 auto poolOp = dyn_cast<T>(op);
426 SmallVector<int64_t> kernelValues;
429 for (
const auto k : kernelValues)
430 if (
failed(levelCheckKernel(op, k,
"kernel")))
434 SmallVector<int64_t> strideValues;
437 for (
const auto s : strideValues)
438 if (
failed(levelCheckStride(op, s,
"stride")))
442 SmallVector<int64_t> padValues;
444 for (
const auto p : padValues)
445 if (
failed(levelCheckKernel(op, p,
"pad")))
453 template <
typename T>
454 LogicalResult levelCheckConv(Operation *op) {
455 if (
auto convOp = dyn_cast<T>(op)) {
457 for (
auto k : convOp.getDilation()) {
458 if (
failed(levelCheckKernel(op, k,
"dilation"))) {
462 for (
auto p : convOp.getPad()) {
463 if (
failed(levelCheckKernel(op, p,
"pad"))) {
467 for (
auto s : convOp.getStride()) {
468 if (
failed(levelCheckStride(op, s,
"stride"))) {
472 auto dilation = convOp.getDilation();
473 if (ShapedType weightType =
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],
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],
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],
508 LogicalResult levelCheckConv2DBlockScaled(Operation *op) {
509 auto convOp = dyn_cast<Conv2DBlockScaledOp>(op);
513 SmallVector<int64_t> padValues;
515 for (
const auto p : padValues)
516 if (
failed(levelCheckKernel(op, p,
"pad <= MAX_KERNEL")))
520 SmallVector<int64_t> strideValues;
523 for (
const auto s : strideValues)
524 if (
failed(levelCheckKernel(op, s,
"stride <= MAX_KERNEL")))
528 SmallVector<int64_t> 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;
540 if (!ShapedType::isDynamic(KH) &&
541 failed(levelCheckKernel(op, dilationValues[0] * KH,
542 "dilation_y * KH <= MAX_KERNEL)")))
545 if (!ShapedType::isDynamic(KW) &&
546 failed(levelCheckKernel(op, dilationValues[1] * KW,
547 "dilation_x * KW <= MAX_KERNEL)")))
555 template <
typename T>
556 LogicalResult levelCheckFFT(Operation *op) {
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"))) {
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);
580 if (
failed(levelCheckKernel(op, shape[1],
"KH")) ||
581 failed(levelCheckKernel(op, shape[2],
"KW"))) {
585 for (
auto p : transpose.getOutPad()) {
586 if (
failed(levelCheckKernel(op, p,
"pad"))) {
590 for (
auto s : transpose.getStride()) {
591 if (
failed(levelCheckStride(op, s,
"stride"))) {
600 LogicalResult levelCheckResize(Operation *op) {
601 if (
auto resize = dyn_cast<tosa::ResizeOp>(op)) {
602 SmallVector<int64_t> scale;
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];
612 levelCheckScale(op, scaleYN / scaleYD,
"scale_y_n/scale_y_d")) ||
614 levelCheckScale(op, scaleXN / scaleXD,
"scale_x_n/scale_x_d"))) {
625 static void getMaxNestedDepth(Operation *op, int32_t &depth) {
626 if (isa<mlir::func::FuncOp>(op) || isa<ModuleOp>(op))
634 getMaxNestedDepth(op, depth);
637 LogicalResult levelCheckMaxNesting(Operation *op) {
638 int32_t maxNestedDepth = 0;
639 getMaxNestedDepth(op, maxNestedDepth);
641 const int32_t maxNestingLevel = targetEnv.getLevel().MAX_NESTING;
642 if (maxNestedDepth >= maxNestingLevel)
644 <<
"failed level check: tosa_nesting_depth < MAX_NESTING" <<
" ("
645 << maxNestingLevel <<
"), got " << maxNestedDepth;
649 LogicalResult levelCheckListSize(Operation *op) {
650 if (
auto concat = dyn_cast<tosa::ConcatOp>(op)) {
651 return levelCheckListSize(op,
concat.getInput1().size(),
"input1");
653 if (
auto custom = dyn_cast<tosa::CustomOp>(op)) {
654 if (
failed(levelCheckListSize(op, custom.getInputList().size(),
656 failed(levelCheckListSize(op, custom.getOutputList().size(),
661 if (
auto condIf = dyn_cast<tosa::IfOp>(op)) {
663 levelCheckListSize(op, condIf.getInputList().size(),
"inputs")) ||
664 failed(levelCheckListSize(op, condIf.getOutputList().size(),
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"))) {
675 if (
auto concat_shape = dyn_cast<tosa::ConcatShapeOp>(op))
676 return levelCheckListSize(op, concat_shape.getInput().size(),
"input");
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)) {
685 <<
"failed attribute check: rounding_mode = DOUBLE_ROUND "
686 <<
"requires extension [doubleround]";
689 if (rescale.getRoundingMode() == RoundingMode::INEXACT_ROUND &&
690 !targetEnv.allows(Extension::inexactround)) {
692 <<
"failed attribute check: rounding_mode = INEXACT_ROUND "
693 <<
"requires extension [inexactround]";
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() &&
707 <<
"failed attribute check: CAST attribute input_unsigned "
708 <<
"requires version 1.1.draft"
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);
724 std::function<LogicalResult(Operation *,
const tosa::TargetEnv &)>>
727 TosaProfileCompliance profileComp;
728 tosa::TargetEnv targetEnv;
732LogicalResult TosaValidation::levelCheckRanks(tosa::ArgMaxOp tosaOp) {
733 auto *op = tosaOp.getOperation();
734 if (
failed(levelCheckRank(op, tosaOp.getInput(),
"operand",
739 if (
failed(levelCheckRank(op, tosaOp.getOutput(),
"result",
747LogicalResult TosaValidation::levelCheckRanks(tosa::ArgMinOp tosaOp) {
748 auto *op = tosaOp.getOperation();
749 if (
failed(levelCheckRank(op, tosaOp.getInput(),
"operand",
754 if (
failed(levelCheckRank(op, tosaOp.getOutput(),
"result",
762LogicalResult TosaValidation::levelCheckRanks(tosa::IfOp tosaOp) {
763 auto *op = tosaOp.getOperation();
766 if (
failed(levelCheckRank(op, tosaOp.getCondition(),
"operand",
774LogicalResult TosaValidation::levelCheckRanks(tosa::VariableOp tosaOp) {
775 auto *op = tosaOp.getOperation();
777 if (
failed(levelCheckRank(op, variableType,
"variable type",
785LogicalResult TosaValidation::levelCheckSizes(tosa::VariableOp tosaOp) {
786 auto *op = tosaOp.getOperation();
788 if (
failed(levelCheckSize(op, variableType,
"variable type")))
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)))) \
799 if (failed(levelCheckSizes(cast<tosa::tosaOp##Op>(op)))) \
803#define CHECK_SIZES(tosaOp) \
804 if (isa<tosa::tosaOp##Op>(op)) { \
805 if (failed(levelCheckSizes(cast<tosa::tosaOp##Op>(op)))) \
809#define CHECK_SHAPE_LEN(tosaOp) \
810 if (isa<tosa::tosaOp##Op>(op)) { \
811 if (failed(levelCheckShapeLengths(cast<tosa::tosaOp##Op>(op)))) \
944#undef CHECK_RANKS_AND_SIZES
946#undef CHECK_SHAPE_LEN
951LogicalResult TosaValidation::levelCheckSize(Operation *op,
952 const Type &typeToCheck,
953 const StringRef operandOrResult) {
954 if (ShapedType type = dyn_cast<ShapedType>(typeToCheck)) {
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);
971 return op->
emitOpError() <<
"failed level check: " << operandOrResult
972 <<
" shape dimension cannot be dynamic when"
973 <<
" targeting TOSA specification version 1.0"
978 int64_t elementBytes = std::max(INT64_C(1), elementBits / 8);
979 int64_t size = elementBytes * type.getNumElements();
986 const int64_t maxSize =
990 <<
"failed level check: " << operandOrResult
991 <<
" tensor size (in bytes) <= (1 << MAX_LOG2_SIZE - 1)";
996LogicalResult TosaValidation::applyLevelCheck(Operation *op) {
1003 if (
failed(levelCheckRanksAndSizes(op)))
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))) {
1021 if (
failed(levelCheckListSize(op))) {
1025 if (isa<tosa::IfOp>(op) || isa<tosa::WhileOp>(op)) {
1026 if (
failed(levelCheckMaxNesting(op))) {
1034LogicalResult TosaValidation::applyAttributeCheck(Operation *op) {
1035 if (
failed(attributeCheckRescale(op)))
1037 if (
failed(attributeCheckCast(op)))
1042inline bool CompatibleTypes(
const mlir::Type &type,
1043 const mlir::Type &declaredType) {
1045 return type == declaredType;
1048LogicalResult TosaValidation::CheckVariable(Operation *op) {
1049 if (
auto variableOp = dyn_cast<mlir::tosa::VariableOp>(op)) {
1050 mlir::StringAttr nameAttr = variableOp.getNameAttr();
1052 if (variablesMap.count(nameAttr))
1053 return op->
emitOpError() <<
"name has already been declared";
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);
1061 variablesMap[nameAttr] = variableType;
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";
1077 auto varType = variablesMap[nameAttr];
1080 auto type = v.getType();
1081 if (!CompatibleTypes(type, varType))
1082 return op->
emitOpError() <<
"operand type does not equal variable type";
1086 auto type = v.getType();
1087 if (!CompatibleTypes(type, varType))
1088 return op->
emitOpError() <<
"result type does not equal variable type";
1095LogicalResult TosaValidation::applyVariableCheck(Operation *op) {
1096 if (
failed(CheckVariable(op)) ||
failed(CheckVariableReadOrWrite(op)))
1101LogicalResult checkErrorIfResize(Operation *op) {
1102 auto resize = dyn_cast<tosa::ResizeOp>(op);
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());
1113 if (!inputType || !outputType)
1114 return op->
emitOpError(
"expect ranked input/output tensor");
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)
1125 "expect input/output height/width dims to be < 16384, ")
1126 <<
"got [OH, OW, IH, IW] = " << sizes;
1129 SmallVector<int64_t> scale;
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];
1139 if (scaleYN > (1 << 11) || scaleXN > (1 << 11))
1141 "expect all scale numerator values to be <= (1 << 11), "
1143 << scaleYN <<
", scale_x_n=" << scaleXN;
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;
1149 SmallVector<int64_t> offset;
1150 SmallVector<int64_t> border;
1155 const int64_t offsetY = offset[0];
1156 const int64_t offsetX = offset[1];
1159 if (offsetY < -scaleYN || offsetY >= 16 * scaleYN)
1161 "expect offsetY / scaleYNumerator to be in range [-1, 16), got ")
1162 << offsetY <<
"/" << scaleYN;
1163 if (offsetX < -scaleXN || offsetX >= 16 * scaleXN)
1165 "expect offsetX / scaleXNumerator to be in range [-1, 16), got ")
1166 << offsetX <<
"/" << scaleXN;
1168 const int64_t borderY = border[0];
1169 const int64_t borderX = border[1];
1170 if (borderY < -16 * scaleYN || borderY >= scaleYN)
1172 "expect borderY / scaleYNumerator to be in range [-16, 1), got ")
1173 << borderY <<
"/" << scaleYN;
1174 if (borderX < -16 * scaleXN || borderX >= scaleXN)
1176 "expect borderX / scaleXNumerator to be in range [-16, 1), got ")
1177 << borderX <<
"/" << scaleXN;
1190 const int64_t
rhs) -> std::optional<int64_t> {
1192 return std::nullopt;
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);
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())
1206 "expected (input_height - 1) * scale_y_n - offset_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)
1214 "calculated output height did not match expected: ")
1215 <<
"calculated=" << calculatedOutHeight <<
", expected=" << oh;
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())
1223 "expected (input_width - 1) * scale_x_n - offset_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;
1237LogicalResult checkErrorIfMul(Operation *op) {
1238 auto mul = dyn_cast<tosa::MulOp>(op);
1244 ElementsAttr shift_elem;
1247 int32_t shift = shift_elem.getValues<IntegerAttr>()[0].getInt();
1249 if (inputElemType.isInteger(32)) {
1251 if (shift < 0 || shift > 63)
1253 <<
"requires 0 <= shift && shift <= 63, but got: " << shift;
1258 <<
"requires shift = 0 for all input data types that "
1259 "are not int32_t, but got: "
1267 auto table = dyn_cast<tosa::TableOp>(op);
1273 const int tableSize = inputElemType.isInteger(8) ? 256 : 513;
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;
1286LogicalResult checkErrorIfRescale(
Operation *op) {
1287 auto rescale = dyn_cast<tosa::RescaleOp>(op);
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())
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();
1305 bool scale32 = rescale.getScale32();
1306 auto roundingMode = rescale.getRoundingMode();
1309 if (scale32 && inWidth == 48)
1310 return op->
emitOpError() <<
"scale32 is not allowed with 48-bit input.";
1313 if (!scale32 && roundingMode == RoundingMode::DOUBLE_ROUND)
1315 <<
"DOUBLE_ROUND is only allowed with scale32=true.";
1318 if (inputUnsigned && outputUnsigned)
1319 return op->
emitOpError() <<
"input and output cannot be both unsigned.";
1322 if (outWidth == 32 && inputUnsigned)
1324 <<
"i32 output type is not allowed with unsigned input.";
1327 if (inWidth == 32 && outputUnsigned)
1329 <<
"i32 input type is not allowed with unsigned output.";
1332 if (inWidth == 48 && outputUnsigned)
1334 <<
"i48 input type is not allowed with unsigned output.";
1337 if (inWidth == 48 && inputUnsigned)
1338 return op->
emitOpError() <<
"i48 input type cannot be unsigned.";
1341 if (inWidth == 32 && inputUnsigned)
1342 return op->
emitOpError() <<
"i32 input type cannot be unsigned.";
1345 if (outWidth == 32 && outputUnsigned)
1346 return op->
emitOpError() <<
"i32 output type cannot be unsigned.";
1351LogicalResult checkErrorIfPad(
Operation *op) {
1352 auto pad = dyn_cast<tosa::PadOp>(op);
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();
1370LogicalResult checkErrorIfReshape(Operation *op) {
1371 auto reshapeOp = dyn_cast<tosa::ReshapeOp>(op);
1375 SmallVector<int64_t> shapeValues;
1381 return op->
emitOpError(
"shape input contains inferable dimension (")
1384 "which does not conform to the TOSA specification";
1389LogicalResult checkErrorIfSlice(Operation *op) {
1390 auto sliceOp = dyn_cast<tosa::SliceOp>(op);
1394 SmallVector<int64_t> startValues;
1395 SmallVector<int64_t> sizeValues;
1397 sliceOp.getStart().getDefiningOp(), startValues);
1398 const bool hasSizeValues =
1402 return op->
emitOpError(
"start input contains inferable dimension (")
1404 <<
") which does not conform to the TOSA specification";
1406 return op->
emitOpError(
"size input contains inferable dimension (")
1409 "does not conform to the TOSA specification";
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);
1421static LogicalResult isRegionIsolatedFromAbove(Region ®ionToCheck) {
1422 bool noLiveInValue =
true;
1423 regionToCheck.
walk([&noLiveInValue, ®ionToCheck](Operation *op) {
1424 if (!isOpIsolatedWithinRegion(op, ®ionToCheck)) {
1425 noLiveInValue =
false;
1430 return noLiveInValue ?
success() : failure();
1433LogicalResult checkIsolatedRegion(Operation *op, Region ®ionToCheck,
1434 StringRef regionName) {
1435 if (succeeded(isRegionIsolatedFromAbove(regionToCheck)))
1438 <<
"is not conformant to the TOSA specification. It requires the '"
1439 << regionName <<
"' region is isolated from above.\n";
1442LogicalResult checkErrorIfCondIf(Operation *op) {
1443 auto ifOp = dyn_cast<tosa::IfOp>(op);
1476 if (
failed(checkIsolatedRegion(op, ifOp.getThenGraph(),
"then")) ||
1477 failed(checkIsolatedRegion(op, ifOp.getElseGraph(),
"else")))
1482LogicalResult checkErrorIfWhileLoop(Operation *op) {
1483 auto whileOp = dyn_cast<tosa::WhileOp>(op);
1487 if (
failed(checkIsolatedRegion(op, whileOp.getCondGraph(),
"cond")) ||
1488 failed(checkIsolatedRegion(op, whileOp.getBodyGraph(),
"body")))
1493LogicalResult checkErrorIfScatter(Operation *op) {
1494 auto scatterOp = dyn_cast<tosa::ScatterOp>(op);
1499 DenseIntElementsAttr indicesAttr;
1503 auto const indicesType =
1504 dyn_cast<ShapedType>(scatterOp.getIndices().getType());
1505 if (!indicesType || !indicesType.hasRank()) {
1511 op->
emitOpError(
"indices values contain duplicates");
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)))
1528LogicalResult TosaValidation::applyFunctionSignatureCheck(func::FuncOp op) {
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";
1541 if (
failed(validateOperationElementTypes(op, !strictOpSpecAlignment)))
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)))
1550 if (
failed(levelCheckSize(op, argType, inputDesc)))
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)))
1557 if (
failed(levelCheckSize(op, resultType, resultDesc)))
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";
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))
1584 }
else if (
auto intTy = dyn_cast<IntegerType>(type)) {
1585 if (intTy.isSignless()) {
1586 switch (intTy.getWidth()) {
1596 }
else if (allowUnsigned && intTy.isUnsigned()) {
1597 switch (intTy.getWidth()) {
1604 }
else if (isa<tosa::shapeType>(type))
1606 else if (isa<tosa::mxint8Type, tosa::BlockScaledType>(type))
1609 return op->
emitOpError() <<
"is not profile-aligned: element type " << type
1614TosaValidation::validateOperationElementTypes(TosaOp op,
bool allowUnsigned) {
1615 for (Value operand : op->getOperands()) {
1617 if (
failed(validateValidElementType(op, elementTy, allowUnsigned)))
1621 for (Type resultTy : op->getResultTypes()) {
1623 if (
failed(validateValidElementType(op, elementTy, allowUnsigned)))
1627 if (
auto variableOp = dyn_cast<tosa::VariableOp>(*op)) {
1629 validateValidElementType(op, variableOp.getType(), allowUnsigned)))
1636TosaValidation::validateOperationElementTypes(func::FuncOp op,
1637 bool allowUnsigned) {
1638 for (
const Type &argType :
1639 llvm::concat<const Type>(op.getArgumentTypes(), op.getResultTypes())) {
1641 if (
failed(validateValidElementType(op, elementTy, allowUnsigned)))
1648void TosaValidation::runOnOperation() {
1649 ModuleOp modOp = getOperation();
1650 TosaDialect *tosaDialect =
getContext().getLoadedDialect<TosaDialect>();
1655 const auto maybeTargetEnv =
1657 if (
failed(maybeTargetEnv))
1658 return signalPassFailure();
1659 targetEnv = *maybeTargetEnv;
1661 const auto functions = modOp.getOps<func::FuncOp>();
1662 if (validateFunctionSignature &&
1663 llvm::any_of(functions, [&](func::FuncOp func) {
1664 return failed(applyFunctionSignatureCheck(func));
1666 return signalPassFailure();
1668 modOp.walk([&](TosaOp op) {
1674 const bool allowUnsigned =
1675 !strictOpSpecAlignment && isa<tosa::RescaleOp>(op);
1676 if (
failed(validateOperationElementTypes(op, allowUnsigned)))
1677 return signalPassFailure();
1679 if (strictOpSpecAlignment &&
1681 return signalPassFailure();
1683 if (strictOpSpecAlignment &&
1685 return signalPassFailure();
1687 if (strictOpSpecAlignment &&
1688 failed(checkSpecificationVersionConstraint(op, targetEnv)))
1689 return signalPassFailure();
1691 if (!allowInvalidOpDatatypeCombinations &&
1693 return signalPassFailure();
1697 if (
failed(applyConstantOperandCheck(op)))
1698 signalPassFailure();
1701 if (
failed(applyLevelCheck(op)))
1702 signalPassFailure();
1705 if (
failed(applyAttributeCheck(op)))
1706 signalPassFailure();
1709 if (
failed(applyVariableCheck(op)))
1710 signalPassFailure();
1713 if (strictOpSpecAlignment &&
failed(applyErrorIfCheck(op)))
1714 signalPassFailure();
static llvm::ManagedStatic< PassManagerOptions > options
static std::optional< int64_t > idivCheck(const int64_t lhs, const int64_t rhs)
#define CHECK_RANKS_AND_SIZES(tosaOp)
#define CHECK_SIZES(tosaOp)
#define CHECK_SHAPE_LEN(tosaOp)
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.
An attribute that represents a reference to a dense integer vector or tensor object.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
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...
Adaptor class to abstract the differences between whether value is from a ShapedType or ShapedTypeCom...
Type getType() const
Return the type of this value.
static WalkResult advance()
static WalkResult interrupt()
This class represents the capability enabled in the target implementation such as profile,...
TosaLevel getLevel() const
static FailureOr< TargetEnv > createTargetEnvFromAttr(TargetEnvAttr targetAttr, Location targetEnvAttrLoc)
bool allows(Profile prof) const
TosaSpecificationVersion getSpecVersion() const
A thin wrapper around the SpecificationVersion enum to represent and provide utilities around the TOS...
bool isBackwardsCompatibleWith(TosaSpecificationVersion baseVersion) const
SmallVector< AffineExpr, 4 > concat(ArrayRef< AffineExpr > a, ArrayRef< AffineExpr > b)
Return the vector that is the concatenation of a and b.
llvm::SmallString< 4 > stringifyVersion(TosaSpecificationVersion version)
RankedTensorType getVariableType(VariableOp variableOp)
static constexpr TosaLevel TOSA_LEVEL_NONE
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...
unsigned getBitWidth(Type type)
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.
@ Mul
RHS of mul is always a constant or a symbolic expression.
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
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.