MLIR 24.0.0git
TosaOps.cpp
Go to the documentation of this file.
1//===- TosaOps.cpp - MLIR Dialect for TOSA --------------------------------===//
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// \file
10// This file implements the TOSA Specification:
11// https://www.mlplatform.org/tosa/tosa_spec.html
12//
13//===----------------------------------------------------------------------===//
14
26#include "mlir/IR/Matchers.h"
30#include "llvm/ADT/APFloat.h"
31#include "llvm/ADT/SmallVectorExtras.h"
32#include "llvm/ADT/TypeSwitch.h"
33
34#include <numeric>
35#include <type_traits>
36
37using namespace mlir;
38using namespace mlir::tosa;
39
40#include "mlir/Dialect/Tosa/IR/TosaOpsDialect.cpp.inc"
42
43//===----------------------------------------------------------------------===//
44// Tosa dialect interface includes.
45//===----------------------------------------------------------------------===//
46
47#include "mlir/Dialect/Tosa/IR/TosaEnums.cpp.inc"
48#include "mlir/Dialect/Tosa/IR/TosaInterfaces.cpp.inc"
49
50namespace {
51#include "mlir/Dialect/Tosa/IR/TosaDialectBytecode.cpp.inc"
52
53//===----------------------------------------------------------------------===//
54// Dialect Function Inliner Interface.
55//===----------------------------------------------------------------------===//
56struct TosaInlinerInterface : public DialectInlinerInterface {
57 using DialectInlinerInterface::DialectInlinerInterface;
58
59 //===--------------------------------------------------------------------===//
60 // Analysis Hooks.
61 //===--------------------------------------------------------------------===//
62
63 /// All operations can be inlined by default.
64 bool isLegalToInline(Operation *op, Region *region, bool wouldBeCloned,
65 IRMapping &map) const final {
66 return true;
67 }
68
69 /// All regions with If and While parent operators can be inlined.
70 bool isLegalToInline(Region *dest, Region *src, bool wouldBeCloned,
71 IRMapping &map) const final {
72 return (isa<tosa::IfOp>(dest->getParentOp()) ||
73 isa<tosa::WhileOp>(dest->getParentOp()));
74 }
75};
76
77/// This class implements the bytecode interface for the Tosa dialect.
78struct TosaDialectBytecodeInterface : public BytecodeDialectInterface {
79 TosaDialectBytecodeInterface(Dialect *dialect)
80 : BytecodeDialectInterface(dialect) {}
81
82 //===--------------------------------------------------------------------===//
83 // Attributes
84
85 Attribute readAttribute(DialectBytecodeReader &reader) const override {
86 return ::readAttribute(getContext(), reader);
87 }
88
89 LogicalResult writeAttribute(Attribute attr,
90 DialectBytecodeWriter &writer) const override {
91 return ::writeAttribute(attr, writer);
92 }
93
94 //===--------------------------------------------------------------------===//
95 // Types
96
97 Type readType(DialectBytecodeReader &reader) const override {
98 return ::readType(getContext(), reader);
99 }
100
101 LogicalResult writeType(Type type,
102 DialectBytecodeWriter &writer) const override {
103 return ::writeType(type, writer);
104 }
105
106 void writeVersion(DialectBytecodeWriter &writer) const final {
107 // TODO: Populate.
108 }
109
110 std::unique_ptr<DialectVersion>
111 readVersion(DialectBytecodeReader &reader) const final {
112 // TODO: Populate
113 reader.emitError("Dialect does not support versioning");
114 return nullptr;
115 }
116
117 LogicalResult upgradeFromVersion(Operation *topLevelOp,
118 const DialectVersion &version) const final {
119 return success();
120 }
121};
122
123} // namespace
124
125//===----------------------------------------------------------------------===//
126// TOSA control flow support.
127//===----------------------------------------------------------------------===//
128
129/// Returns the while loop body.
130SmallVector<Region *> tosa::WhileOp::getLoopRegions() {
131 return {&getBodyGraph()};
132}
133
134//===----------------------------------------------------------------------===//
135// TOSA variable operator support.
136//===----------------------------------------------------------------------===//
137
139 return map_to_vector(shape, [](int64_t dim) {
140 return dim == -1 ? ShapedType::kDynamic : dim;
141 });
142}
143
144// returns type of variable op
145RankedTensorType mlir::tosa::getVariableType(tosa::VariableOp variableOp) {
146 Type elementType = variableOp.getType();
147 DenseIntElementsAttr varShapeAttr = variableOp.getVarShape();
148 auto shape = convertToMlirShape(to_vector(varShapeAttr.getValues<int64_t>()));
149 return RankedTensorType::get(shape, elementType);
150}
151
152//===----------------------------------------------------------------------===//
153// Tosa dialect initialization.
154//===----------------------------------------------------------------------===//
155
156void TosaDialect::initialize() {
157 addTypes<
158#define GET_TYPEDEF_LIST
159#include "mlir/Dialect/Tosa/IR/TosaOpsTypesBase.cpp.inc"
160 >();
161 addOperations<
162#define GET_OP_LIST
163#include "mlir/Dialect/Tosa/IR/TosaOps.cpp.inc"
164 >();
165 addAttributes<
166#define GET_ATTRDEF_LIST
167#include "mlir/Dialect/Tosa/IR/TosaAttributes.cpp.inc"
168 >();
169 addInterfaces<TosaDialectBytecodeInterface, TosaInlinerInterface>();
170 declarePromisedInterfaces<
171 shard::ShardingInterface, ClampOp, SigmoidOp, TanhOp, AddOp,
172 ArithmeticRightShiftOp, BitwiseAndOp, BitwiseOrOp, BitwiseXorOp, IntDivOp,
173 LogicalAndOp, LogicalLeftShiftOp, LogicalRightShiftOp, LogicalOrOp,
174 LogicalXorOp, MaximumOp, MinimumOp, MulOp, PowOp, SubOp, AbsOp,
175 BitwiseNotOp, CeilOp, ClzOp, ExpOp, FloorOp, LogOp, LogicalNotOp,
176 NegateOp, ReciprocalOp, RsqrtOp, SelectOp, EqualOp, GreaterOp,
177 GreaterEqualOp, MatMulOp>();
178}
179
180Operation *TosaDialect::materializeConstant(OpBuilder &builder, Attribute value,
181 Type type, Location loc) {
182 // Tosa dialect constants only support ElementsAttr unlike standard dialect
183 // constant which supports all attributes.
184 if (llvm::isa<shapeType>(type) && llvm::isa<DenseIntElementsAttr>(value)) {
185 return tosa::ConstShapeOp::create(builder, loc, type,
186 llvm::cast<DenseIntElementsAttr>(value));
187 }
188 if (llvm::isa<ElementsAttr>(value))
189 return tosa::ConstOp::create(builder, loc, type,
190 llvm::cast<ElementsAttr>(value));
191 return nullptr;
192}
193
194//===----------------------------------------------------------------------===//
195// Parsers and printers
196//===----------------------------------------------------------------------===//
197
198namespace {
199
200ParseResult getShapeAndElementType(OpAsmParser &parser, Type parsedType,
201 DenseElementsAttr &varShapeAttr,
202 TypeAttr &typeAttr) {
203 if (auto shapedType = dyn_cast<ShapedType>(parsedType)) {
204 if (!shapedType.hasRank())
205 return parser.emitError(parser.getCurrentLocation())
206 << "expected ranked type";
207
208 auto elementType = shapedType.getElementType();
209 typeAttr = TypeAttr::get(elementType);
210 ArrayRef<int64_t> shape = shapedType.getShape();
211 Builder builder(parser.getContext());
212 varShapeAttr = builder.getIndexTensorAttr(convertFromMlirShape(shape));
213 return success();
214 }
215 return parser.emitError(parser.getCurrentLocation())
216 << "expected shaped type";
217}
218
219} // namespace
220
221// parses the optional initial value or type for a tosa variable
222// with initial value:
223// tosa.variable @name = dense<0.0> : tensor<1x8xf32>
224//
225// without initial value:
226// tosa.variable @name : tensor<1x8xf32>
228 OpAsmParser &parser, DenseElementsAttr &varShapeAttr, TypeAttr &typeAttr,
229 Attribute &initialValueAttr) {
230 if (succeeded(parser.parseOptionalEqual())) {
231 if (failed(parser.parseAttribute(initialValueAttr))) {
232 return parser.emitError(parser.getCurrentLocation())
233 << "expected attribute";
234 }
235 if (auto typedAttr = dyn_cast<TypedAttr>(initialValueAttr)) {
236 return getShapeAndElementType(parser, typedAttr.getType(), varShapeAttr,
237 typeAttr);
238 }
239 return parser.emitError(parser.getCurrentLocation())
240 << "expected Typed attr";
241 }
242
243 initialValueAttr = nullptr;
244 Type parsedType;
245 if (failed(parser.parseColonType(parsedType))) {
246 return parser.emitError(parser.getCurrentLocation())
247 << "expected type after colon";
248 }
249 return getShapeAndElementType(parser, parsedType, varShapeAttr, typeAttr);
250}
251
253 OpAsmPrinter &p, Operation *op, DenseElementsAttr varShapeAttr,
254 TypeAttr typeAttr, Attribute initialValueAttr) {
255 bool needsSpace = false;
256 if (!dyn_cast_or_null<TypedAttr>(initialValueAttr)) {
257 auto shape =
258 convertToMlirShape(to_vector(varShapeAttr.getValues<int64_t>()));
259 Type elementType = typeAttr.getValue();
260 RankedTensorType tensorType =
261 RankedTensorType::get(ArrayRef<int64_t>(shape), elementType);
262 auto tensorTypeAttr = TypeAttr::get(tensorType);
263 p << ": ";
264 p.printAttribute(tensorTypeAttr);
265 needsSpace = true; // subsequent attr value needs a space separator
266 }
267 if (initialValueAttr) {
268 if (needsSpace)
269 p << ' ';
270 p << "= ";
271 p.printAttribute(initialValueAttr);
272 }
273}
274
275//===----------------------------------------------------------------------===//
276// Tosa utilities.
277//===----------------------------------------------------------------------===//
278
279static std::optional<int64_t> idivCheck(const int64_t lhs, const int64_t rhs) {
280 if (lhs % rhs != 0)
281 return std::nullopt;
282 return lhs / rhs;
283}
284
286 auto srcType = getElementTypeOrSelf(type);
287 if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(srcType))
288 srcType = getStorageElementTypeFromQuantized(quantType);
289 return srcType;
290}
291
295
296static LogicalResult verifyRescaleValueAndZpTypes(Operation *op, Value val,
297 Value valZp, StringRef name) {
299 Type eZpType = getStorageElementTypeOrSelf(valZp.getType());
300
301 bool bothInts =
302 mlir::isa<IntegerType>(eType) && mlir::isa<IntegerType>(eZpType);
303 bool sameBitWidth =
304 (eType.getIntOrFloatBitWidth() == eZpType.getIntOrFloatBitWidth());
305
306 if (!bothInts || !sameBitWidth) {
307 return op->emitOpError()
308 << "expected " << name << " and " << name
309 << "_zp to both be integer of the same bitwidth, but got " << eType
310 << " vs. " << eZpType;
311 }
312 return success();
313}
314
315// Create a pad-const const tensor with value of `val` of required data-type
317 Value src, int32_t val) {
318 const auto srcType = getElementTypeOrSelf(src);
319 const auto srcElemType = getStorageElementTypeOrSelf(src);
320 const auto padConstType = mlir::RankedTensorType::get({1}, srcType);
321 const auto padConstEType = mlir::RankedTensorType::get({1}, srcElemType);
322 const auto padConstAttr{
323 llvm::isa<FloatType>(srcElemType)
324 ? DenseElementsAttr::get(padConstEType,
325 builder.getFloatAttr(srcElemType, val))
326 : DenseElementsAttr::get(padConstEType,
327 builder.getIntegerAttr(srcElemType, val))};
328 return tosa::ConstOp::create(builder, loc, padConstType, padConstAttr);
329}
330
332 if (auto blockScaledTy = dyn_cast<tosa::BlockScaledType>(type))
333 return getBitWidth(blockScaledTy.getValueType());
334 if (dyn_cast<tosa::mxint8Type>(type))
335 return 8;
336 return type.getIntOrFloatBitWidth();
337}
338
339// Update dim size if current dim is dynamic, otherwise raise an error if sizes
340// do not match
341LogicalResult tryUpdateDimOrFailure(Operation *op, int64_t &currDim,
342 const int64_t newDim,
343 const StringRef operandName,
344 const StringRef dimName) {
345 if (ShapedType::isDynamic(currDim)) {
346 currDim = newDim;
347 return success();
348 } else if (ShapedType::isStatic(newDim) && currDim != newDim) {
349 return op->emitOpError("expected ")
350 << dimName << " of " << operandName << " to match size " << currDim
351 << ", got " << newDim;
352 }
353 return success();
354}
355
358 auto printDim = [&](int64_t dim) {
359 if (ShapedType::isDynamic(dim))
360 diag << "?";
361 else
362 diag << dim;
363 };
364
365 llvm::interleaveComma(shape, diag, printDim);
366}
367
368static LogicalResult
370 ArrayRef<int64_t> expectedShape,
371 StringRef outputName = "output") {
372 assert(outputType.hasRank() && "expected output type to be ranked");
373
374 if (succeeded(verifyCompatibleShape(outputType.getShape(), expectedShape)))
375 return success();
376
377 InFlightDiagnostic diag = op->emitOpError("expected ");
378 diag << outputName << " shape ";
379 printShapeToDiagnostic(diag, outputType.getShape());
380 diag << " to be compatible with inferred shape ";
381 printShapeToDiagnostic(diag, expectedShape);
382 return diag;
383}
384
386 Operation *op, const int64_t inputSize, const int64_t kernelSize,
387 const int64_t outputSize, const int64_t padBefore, const int64_t padAfter,
388 const int64_t stride, const int64_t dilation, const llvm::StringRef dimName,
389 const llvm::StringRef dimAxis, const llvm::StringRef padBeforeName,
390 const llvm::StringRef padAfterName) {
391 if (inputSize == ShapedType::kDynamic || kernelSize == ShapedType::kDynamic)
392 return success();
393
394 // ERROR_IF: O != idiv_check(I - 1 + pa + pb - (K - 1) * d, s) + 1
395
396 const std::optional<int64_t> calculatedOutSizeMinusOne = idivCheck(
397 inputSize - 1 + padBefore + padAfter - (kernelSize - 1) * dilation,
398 stride);
399 if (!calculatedOutSizeMinusOne.has_value())
400 return op->emitOpError("expected input_")
401 << dimName << " - 1 + pad_" << padBeforeName << " + pad_"
402 << padAfterName << " - (kernel_" << dimName << " - 1) * dilation_"
403 << dimAxis << " to be wholly divisible by stride_" << dimAxis
404 << ", got (" << inputSize << " - 1 + " << padBefore << " + "
405 << padAfter << " - (" << kernelSize << " - 1) * " << dilation
406 << ") / " << stride;
407
408 const int64_t calculatedOutSize = calculatedOutSizeMinusOne.value() + 1;
409 if (outputSize != ShapedType::kDynamic && calculatedOutSize != outputSize)
410 return op->emitOpError("calculated output ")
411 << dimName << " did not match expected: "
412 << "calculated=" << calculatedOutSize << ", expected=" << outputSize;
413
414 return success();
415}
416
417//===----------------------------------------------------------------------===//
418// mxint8Type DenseElementTypeInterface implementation.
419//===----------------------------------------------------------------------===//
420size_t mlir::tosa::mxint8Type::getDenseElementBitSize() const { return 8; }
421
423mlir::tosa::mxint8Type::convertToAttribute(ArrayRef<char> rawData) const {
424 assert(rawData.size() == 1 && "expected 1 byte for tosa.mxint8 element");
425 const auto intType = IntegerType::get(getContext(), 8);
426 return intType.convertToAttribute(rawData);
427}
428
429LogicalResult mlir::tosa::mxint8Type::convertFromAttribute(
431 const auto intAttr = dyn_cast<IntegerAttr>(attr);
432 if (!intAttr)
433 return failure();
434 const Type attrType = intAttr.getType();
435 if (!attrType.isSignlessInteger(8))
436 return failure();
437 return cast<IntegerType>(attrType).convertFromAttribute(attr, result);
438}
439
440//===----------------------------------------------------------------------===//
441// TOSA block scaling utilities.
442//===----------------------------------------------------------------------===//
443
446 bool allowScaleValues) {
447 const auto tensorType = llvm::cast<ShapedType>(type);
448 const BlockScaledType elemType =
449 llvm::dyn_cast<BlockScaledType>(tensorType.getElementType());
450 if (!elemType)
451 return success();
452
453 if (!allowScaleValues && elemType.hasScaleValues()) {
454 if (emitError)
455 emitError()
456 << "block scaled tensor type with scale values is not allowed";
457 return failure();
458 }
459
460 if (!tensorType.hasRank())
461 return success();
462
463 if (tensorType.getRank() == 0) {
464 if (emitError)
465 emitError() << "block scaled tensor type must have rank greater than "
466 "zero";
467 return failure();
468 }
469
470 const ArrayRef<int64_t> tensorShape = tensorType.getShape();
471 const uint32_t blockSize =
472 BlockShapeAttr::getBlockShapeValue(elemType.getBlockShape());
473
474 if (allowScaleValues && elemType.hasScaleValues() &&
475 tensorType.hasStaticShape()) {
476 const size_t numBlocks = tensorType.getNumElements() / blockSize;
477 if (elemType.getScaleValues().size() != numBlocks) {
478 if (emitError)
479 emitError() << "block scaled tensor type with scale values must have "
480 "scale values for each block, expected "
481 << numBlocks << ", got "
482 << elemType.getScaleValues().size();
483 return failure();
484 }
485 }
486
487 const int64_t blockedDimension = tensorShape.back();
488 if (ShapedType::isDynamic(blockedDimension))
489 return success();
490
491 if (blockedDimension % blockSize != 0) {
492 if (emitError)
493 emitError() << "last dimension of block scaled tensor type ("
494 << blockedDimension << ") must be divisible by block size ("
495 << blockSize << ")";
496
497 return failure();
498 }
499
500 return success();
501}
502
504 MLIRContext *ctx = type.getContext();
505 std::string message;
507 ctx, [&](Diagnostic &diag) { message = diag.str(); });
508
510 type, [ctx] { return emitError(UnknownLoc::get(ctx)); })) &&
511 !message.empty()) {
512 return ": " + message;
513 }
514
515 return "";
516}
517
518static ParseResult parseScaleValues(AsmParser &parser,
519 SmallVector<Attribute> &scaleValues,
520 Type scaleType) {
521 const auto parseScaleValue = [&]() -> ParseResult {
522 const SMLoc loc = parser.getCurrentLocation();
523
524 double floatValue;
525 if (parser.parseFloat(floatValue))
526 return failure();
527
528 if (floatValue < 0.0)
529 return parser.emitError(loc, "scale value must be non-negative, got ")
530 << floatValue;
531
532 Type attrType = scaleType;
533 if (succeeded(parser.parseOptionalColon()) && parser.parseType(attrType))
534 return failure();
535
536 if (attrType != scaleType)
537 return parser.emitError(loc, "parsed attribute type ")
538 << attrType << " does not match expected scale type " << scaleType;
539
540 scaleValues.push_back(FloatAttr::get(attrType, floatValue));
541 return success();
542 };
543
544 return parser.parseCommaSeparatedList(parseScaleValue);
545}
546
547static void printScaleValues(AsmPrinter &printer,
548 ArrayRef<Attribute> scaleValues, Type) {
549 llvm::interleaveComma(scaleValues, printer, [&](Attribute scaleValue) {
550 printer.printAttributeWithoutType(scaleValue);
551 });
552}
553
554size_t mlir::tosa::BlockScaledType::getDenseElementBitSize() const {
555 const Type valueType = getValueType();
556 if (isa<tosa::mxint8Type>(valueType))
557 return 8;
558 return valueType.getIntOrFloatBitWidth();
559}
560
562mlir::tosa::BlockScaledType::convertToAttribute(ArrayRef<char> rawData) const {
563 // Block scaled values are stored as a single byte. This is because possible
564 // value data types are either 8-bit or sub-byte. Sub-byte types are aligned
565 // to 8-bits.
566 assert(rawData.size() == 1 && "expected 1 byte for block_scaled element");
567 const Type valueType = getValueType();
568 if (const auto mxint8Value = dyn_cast<tosa::mxint8Type>(valueType))
569 return mxint8Value.convertToAttribute(rawData);
570 if (!isa<FloatType>(valueType))
571 return {};
572 return mlir::detail::convertFloatTypeToAttribute(valueType, rawData);
573}
574
575LogicalResult mlir::tosa::BlockScaledType::convertFromAttribute(
577 const Type valueType = getValueType();
578 if (const auto mxint8Value = dyn_cast<tosa::mxint8Type>(valueType))
579 return mxint8Value.convertFromAttribute(attr, result);
580
581 const auto floatAttr = dyn_cast<FloatAttr>(attr);
582 if (!floatAttr || floatAttr.getType() != valueType)
583 return failure();
584 // const APFloat value = floatAttr.getValue();
585 return mlir::detail::convertFloatTypeFromAttribute(valueType, floatAttr,
586 result);
587}
588
589//===----------------------------------------------------------------------===//
590// TOSA Operator shape inference
591//===----------------------------------------------------------------------===//
592
593template <typename A, std::enable_if_t<std::is_same_v<A, ArgMaxOp::Adaptor> ||
594 std::is_same_v<A, ArgMinOp::Adaptor>,
595 int> = 0>
597 MLIRContext *context, ::std::optional<Location> location, A adaptor,
598 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
599 ShapeAdaptor inputShape(adaptor.getInput().getType());
600 IntegerAttr axis = adaptor.getProperties().axis;
601 int32_t axisVal = axis.getValue().getSExtValue();
602
603 if (!inputShape.hasRank()) {
604 inferredReturnShapes.push_back(ShapedTypeComponents());
605 return success();
606 }
607
608 const auto inputRank = inputShape.getRank();
609 SmallVector<int64_t> outShape;
610 outShape.reserve(inputRank - 1);
611 for (int i = 0, s = inputRank; i < s; i++) {
612 if (i == axisVal)
613 continue;
614 outShape.push_back(inputShape.getDimSize(i));
615 }
616
617 inferredReturnShapes.push_back(ShapedTypeComponents(outShape));
618 return success();
619}
620
621//===----------------------------------------------------------------------===//
622// TOSA Operator Verifiers.
623//===----------------------------------------------------------------------===//
624template <typename T>
625LogicalResult argMaxMinVerify(T op) {
626 const ShapedType resultType = llvm::cast<ShapedType>(op.getType());
627
628 if (const auto resultETy = resultType.getElementType();
629 !resultETy.isIntOrIndex())
630 return op.emitOpError("result tensor is not of integer type");
631
632 const auto inputType = llvm::cast<ShapedType>(op.getInput().getType());
633 if (!inputType.hasRank())
634 return success();
635
636 // Ensure axis is within the tensor rank
637 const int64_t axis = op.getAxisAttr().getInt();
638 if (((axis < 0) || axis >= inputType.getRank()))
639 return op.emitOpError("specified axis is outside the rank of the tensor");
640
641 if (!resultType.hasRank())
642 return success();
643
644 const ArrayRef<int64_t> inputShape = inputType.getShape();
645 const ArrayRef<int64_t> outputShape = resultType.getShape();
646 llvm::SmallVector<int64_t> expectedOutputShape(inputShape);
647 expectedOutputShape.erase(expectedOutputShape.begin() + axis);
648 if (failed(verifyCompatibleShape(expectedOutputShape, outputShape)))
649 return op.emitOpError("expected output shape '")
650 << expectedOutputShape << "', got '" << outputShape << "'";
651
652 return success();
653}
654
655template <typename T>
656static LogicalResult verifyConvOp(T op) {
657 const auto inputType = llvm::dyn_cast<TensorType>(op.getInput().getType());
658 const auto weightType = llvm::dyn_cast<TensorType>(op.getWeight().getType());
659
660 auto inputEType = inputType.getElementType();
661 auto weightEType = weightType.getElementType();
662 auto biasEType =
663 llvm::cast<ShapedType>(op.getBias().getType()).getElementType();
664 auto resultEType =
665 llvm::cast<ShapedType>(op.getResult().getType()).getElementType();
666 bool biasIsFloat = llvm::isa<FloatType>(biasEType);
667 bool resultIsFloat = llvm::isa<FloatType>(resultEType);
668
669 if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(inputEType))
670 inputEType = getStorageElementTypeFromQuantized(quantType);
671
672 if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(weightEType))
673 weightEType = getStorageElementTypeFromQuantized(quantType);
674
675 if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(biasEType))
676 biasEType = getStorageElementTypeFromQuantized(quantType);
677
678 if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(resultEType))
679 resultEType = getStorageElementTypeFromQuantized(quantType);
680
681 if (biasIsFloat && resultIsFloat && (biasEType != resultEType)) {
682 // for now, only enforce bias element type == result element type for
683 // float types.
684 op.emitOpError(
685 "expect both bias and result to have same element type, got ")
686 << biasEType << " and " << resultEType;
687 return failure();
688 }
689
690 const bool isInputBlockScaled = llvm::isa<BlockScaledType>(inputEType);
691 const bool isWeightBlockScaled = llvm::isa<BlockScaledType>(weightEType);
692 const bool isInputFloat = llvm::isa<FloatType>(inputEType);
693 const bool isWeightFloat = llvm::isa<FloatType>(weightEType);
694
695 const bool isInputBSorFloat = isInputBlockScaled || isInputFloat;
696 const bool isWeightBSorFloat = isWeightBlockScaled || isWeightFloat;
697
698 // Either both must be float or both non-float.
699 if (isInputBSorFloat != isWeightBSorFloat) {
700 op.emitOpError(
701 "expect both input and weight to be float or not together, got ")
702 << inputEType << " and " << weightEType;
703 return failure();
704 }
705
706 auto inputZpEType = getStorageElementTypeOrSelf(op.getInputZp().getType());
707 if (!isInputBlockScaled && inputEType != inputZpEType) {
708 return op.emitOpError("expect both input and its zero point are the same "
709 "element type, got ")
710 << inputEType << " and " << inputZpEType;
711 }
712 if (isInputBlockScaled && !llvm::isa<Float32Type>(inputZpEType)) {
713 return op.emitOpError(
714 "expect block scaled input to have fp32 zero point, got ")
715 << inputEType << " and " << inputZpEType;
716 }
717
718 auto weightZpEType = getStorageElementTypeOrSelf(op.getWeightZp().getType());
719 if (!isWeightBlockScaled && weightEType != weightZpEType) {
720 return op.emitOpError("expect both weight and its zero point are the same "
721 "element type, got ")
722 << weightEType << " and " << weightZpEType;
723 }
724 if (isWeightBlockScaled && !llvm::isa<Float32Type>(weightZpEType)) {
725 return op.emitOpError(
726 "expect block scaled weight to have fp32 zero point, got ")
727 << weightEType << " and " << weightZpEType;
728 }
729
730 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
731 if (succeeded(maybeIZp) && op.verifyInputZeroPoint(*maybeIZp).failed())
732 return failure();
733
734 FailureOr<int64_t> maybeWZp = op.getWeightZeroPoint();
735 if (succeeded(maybeWZp) && op.verifyWeightZeroPoint(*maybeWZp).failed())
736 return failure();
737
738 return success();
739}
740
741LogicalResult tosa::ConstOp::verify() {
742 Operation &op = *getOperation();
743 auto attrType = llvm::dyn_cast<TensorType>(getValuesAttr().getType());
744 auto outputType = llvm::dyn_cast<TensorType>(getOutput().getType());
745
746 if (!attrType || !outputType) {
747 emitOpError("expected tensors for attr/result type");
748 return failure();
749 }
750
751 const Type attrElemType = attrType.getElementType();
752 const Type resultElemType = outputType.getElementType();
753
754 if (auto result =
755 llvm::dyn_cast<mlir::quant::QuantizedType>(resultElemType)) {
756 if (getStorageElementTypeFromQuantized(result) == attrElemType)
757 return success();
758 }
759
760 if (auto attrBlockScaledType =
761 llvm::dyn_cast<mlir::tosa::BlockScaledType>(attrElemType)) {
762 if (!attrBlockScaledType.hasScaleValues())
763 return op.emitOpError(
764 "attribute block scaled type must have scale values");
765
766 const auto emitAttributeError = [&op]() {
767 return op.emitOpError("attribute block scaled type is invalid: ");
768 };
769
770 if (failed(verifyBlockScaledTensorType(attrType, emitAttributeError, true)))
771 return failure();
772
773 const BlockScaledType resultBlockScaledType =
774 llvm::dyn_cast<mlir::tosa::BlockScaledType>(resultElemType);
775 if (!resultBlockScaledType)
776 return op.emitOpError(
777 "result type must be block scaled type if attribute is block "
778 "scaled type");
779
780 if (attrBlockScaledType.getValueType() !=
781 resultBlockScaledType.getValueType() ||
782 attrBlockScaledType.getScaleType() !=
783 resultBlockScaledType.getScaleType() ||
784 attrBlockScaledType.getBlockShape() !=
785 resultBlockScaledType.getBlockShape())
786 return op.emitOpError(
787 "expected block scaled element type to be compatible "
788 "between attr and result, got ")
789 << attrBlockScaledType << " vs. " << resultBlockScaledType;
790
791 return success();
792 }
793
794 if (attrElemType != resultElemType)
795 return emitOpError("expected same attr/result element types");
796
797 return success();
798}
799
800template <typename T>
801static LogicalResult verifyConvOpModes(T op) {
802 auto inputEType =
803 llvm::cast<ShapedType>(op.getInput().getType()).getElementType();
804
805 if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(inputEType))
806 inputEType = getStorageElementTypeFromQuantized(quantType);
807
808 auto resultEType =
809 llvm::cast<ShapedType>(op.getResult().getType()).getElementType();
810
811 if (auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(resultEType))
812 resultEType = getStorageElementTypeFromQuantized(quantType);
813
814 return success();
815}
816
817//===----------------------------------------------------------------------===//
818// ERROR_IF functions.
819// ERROR_IF is a predicate that must set an error if the condition holds.
820//===----------------------------------------------------------------------===//
821
822template <typename T>
823static LogicalResult verifyConvOpErrorIf(T op) {
824 llvm::ArrayRef<int64_t> padding = op.getPad();
825 if (llvm::any_of(padding, [](int64_t p) { return p < 0; }))
826 return op.emitOpError("expect all padding values to be >= 0, got ")
827 << padding;
828
829 llvm::ArrayRef<int64_t> strides = op.getStride();
830 if (llvm::any_of(strides, [](int64_t s) { return s < 1; }))
831 return op.emitOpError("expect all stride values to be >= 1, got ")
832 << strides;
833
834 llvm::ArrayRef<int64_t> dilations = op.getDilation();
835 if (llvm::any_of(dilations, [](int64_t d) { return d < 1; }))
836 return op.emitOpError("expect all dilation values to be >= 1, got ")
837 << dilations;
838
839 const RankedTensorType outputType =
840 llvm::dyn_cast<RankedTensorType>(op.getOutput().getType());
841 if (!outputType)
842 // Skip following checks if output is not ranked
843 return success();
844
845 const RankedTensorType inputType =
846 llvm::dyn_cast<RankedTensorType>(op.getInput().getType());
847 const RankedTensorType weightType =
848 llvm::dyn_cast<RankedTensorType>(op.getWeight().getType());
849
850 if (inputType && weightType) {
851 // input = [_,IH,IW,_], weight = [_,KH,KW,_], output = [_,OH,OW,_]
852 if constexpr (std::is_same<T, tosa::Conv2DOp>::value) {
853 if (failed(verifyConvOutputSize(
854 op, inputType.getDimSize(1), weightType.getDimSize(1),
855 outputType.getDimSize(1), padding[0], padding[1], strides[0],
856 dilations[0], "height", "y", "top", "bottom")))
857 return failure();
858
859 if (failed(verifyConvOutputSize(
860 op, inputType.getDimSize(2), weightType.getDimSize(2),
861 outputType.getDimSize(2), padding[2], padding[3], strides[1],
862 dilations[1], "width", "x", "left", "right")))
863 return failure();
864 }
865
866 // input = [_,IH,IW,_], weight = [KH,KW,_,_], output = [_,OH,OW,_]
867 if constexpr (std::is_same<T, tosa::DepthwiseConv2DOp>::value) {
868 if (failed(verifyConvOutputSize(
869 op, inputType.getDimSize(1), weightType.getDimSize(0),
870 outputType.getDimSize(1), padding[0], padding[1], strides[0],
871 dilations[0], "height", "y", "top", "bottom")))
872 return failure();
873
874 if (failed(verifyConvOutputSize(
875 op, inputType.getDimSize(2), weightType.getDimSize(1),
876 outputType.getDimSize(2), padding[2], padding[3], strides[1],
877 dilations[1], "width", "x", "left", "right")))
878 return failure();
879 }
880
881 // input = [_,ID,IH,IW,_], weight = [_,KD,KH,KW,_], output = [_,OD,OH,OW,_]
882 if constexpr (std::is_same<T, tosa::Conv3DOp>::value) {
883 if (failed(verifyConvOutputSize(
884 op, inputType.getDimSize(1), weightType.getDimSize(1),
885 outputType.getDimSize(1), padding[0], padding[1], strides[0],
886 dilations[0], "depth", "d", "front", "back")))
887 return failure();
888
889 if (failed(verifyConvOutputSize(
890 op, inputType.getDimSize(2), weightType.getDimSize(2),
891 outputType.getDimSize(2), padding[2], padding[3], strides[1],
892 dilations[1], "height", "y", "top", "bottom")))
893 return failure();
894
895 if (failed(verifyConvOutputSize(
896 op, inputType.getDimSize(3), weightType.getDimSize(3),
897 outputType.getDimSize(3), padding[4], padding[5], strides[2],
898 dilations[2], "width", "x", "left", "right")))
899 return failure();
900 }
901 }
902
903 const RankedTensorType biasType =
904 llvm::dyn_cast<RankedTensorType>(op.getBias().getType());
905 if (!biasType)
906 // Skip following checks if bias is not ranked
907 return success();
908
909 const int64_t biasChannels = biasType.getDimSize(0);
910 const int64_t outputChannels =
911 outputType.getDimSize(outputType.getRank() - 1);
912 if (biasChannels == ShapedType::kDynamic ||
913 outputChannels == ShapedType::kDynamic)
914 // Skip following checks if biasChannels or outputChannels is dynamic dim
915 return success();
916
917 if (biasChannels != outputChannels && biasChannels != 1)
918 return op.emitOpError(
919 "bias channels expected to be equal to output channels (")
920 << outputChannels << ") or 1, got " << biasChannels;
921
922 return success();
923}
924
925// Verify whether same type and shape of the given two types.
926static LogicalResult errorIfTypeOrShapeMismatch(Operation *op, Type type1,
927 StringRef name1, Type type2,
928 StringRef name2) {
929 auto shapeType1 = dyn_cast<ShapedType>(type1);
930 auto shapeType2 = dyn_cast<ShapedType>(type2);
931 if (!shapeType1 || !shapeType2)
932 return failure();
933
934 auto elemType1 = shapeType1.getElementType();
935 auto elemType2 = shapeType2.getElementType();
936 if (elemType1 != elemType2)
937 return op->emitOpError()
938 << "require same element type for " << name1 << " (" << elemType1
939 << ") and " << name2 << " (" << elemType2 << ")";
940
941 if (failed(verifyCompatibleShape(type1, type2)))
942 return op->emitOpError()
943 << "require same shapes for " << name1 << " (" << type1 << ") and "
944 << name2 << " (" << type2 << ")";
945
946 return success();
947}
948
949// Verify whether same length, type, and shape of the given two tensor lists.
950static LogicalResult errorIfTypeOrShapeMismatch(Operation *op, ValueRange list1,
951 StringRef name1,
952 ValueRange list2,
953 StringRef name2) {
954 if (list1.size() != list2.size())
955 return op->emitOpError()
956 << "require same number of values in " << name1 << " ("
957 << list1.size() << ") and " << name2 << " (" << list2.size() << ")";
958
959 for (auto [type1, type2] :
960 llvm::zip_equal(list1.getTypes(), list2.getTypes())) {
961 if (errorIfTypeOrShapeMismatch(op, type1, name1, type2, name2).failed())
962 return failure();
963 }
964
965 return success();
966}
967
968static inline LogicalResult errorIfShapeNotSizeOne(Operation *op, Type type) {
969 ShapeAdaptor shapeAdaptor(type);
970 if (!shapeAdaptor.hasRank() || !shapeAdaptor.hasStaticShape())
971 return success();
972
973 return shapeAdaptor.getNumElements() == 1 ? success() : failure();
974}
975
976template <typename T>
977static LogicalResult verifyVariableOpErrorIf(T op, Type type, StringRef name) {
978 Operation *symTableOp =
979 op->template getParentWithTrait<OpTrait::SymbolTable>();
980 if (!symTableOp)
981 // If the operation is not the scope of a symbol table, we cannot
982 // verify it against it's declaration.
983 return success();
984
985 SymbolTable symTable(symTableOp);
986 const auto varOp = symTable.lookup<tosa::VariableOp>(op.getName());
987
988 // Verify prior declaration
989 if (!varOp)
990 return op->emitOpError("'")
991 << op.getName() << "' has not been declared by 'tosa.variable'";
992
993 // Verify type and shape
994 auto variableType = getVariableType(varOp);
995 if (errorIfTypeOrShapeMismatch(op, type, name, variableType,
996 "the input tensor")
997 .failed())
998 return failure();
999 return success();
1000}
1001
1002// verify that inType and outType have same element types
1003static LogicalResult verifySameElementTypes(Operation *op, Type aType,
1004 Type bType,
1005 StringRef aName = "input",
1006 StringRef bName = "output") {
1007 auto aTType = llvm::dyn_cast<TensorType>(aType);
1008 auto bTType = llvm::dyn_cast<TensorType>(bType);
1009 if (!aTType) {
1010 op->emitOpError("expect shaped tensor for") << aName << ", got " << aType;
1011 return failure();
1012 }
1013 if (!bTType) {
1014 op->emitOpError("expect shaped tensor for") << bName << ", got" << bType;
1015 return failure();
1016 }
1017 auto aElementType = aTType.getElementType();
1018 auto bElementType = bTType.getElementType();
1019 auto aQuantType =
1020 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(aElementType);
1021 auto bQuantType =
1022 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(bElementType);
1023 if ((aElementType.isIntOrIndexOrFloat() || aQuantType) &&
1024 (bElementType.isIntOrIndexOrFloat() || bQuantType) &&
1025 aElementType != bElementType) {
1026 // only check if both element types are int/index/float/UniformQuantized
1027 // eg, not sure how to check quant::QuantizedType
1028 // this happens in test_conv2d_q_grouped_convolution in
1029 // tfl-to-tosa-pipeline.mlir
1030 op->emitOpError("expect ")
1031 << aName << " and " << bName << " to have same element type, got "
1032 << aElementType << " and " << bElementType;
1033 return failure();
1034 }
1035 return success();
1036}
1037
1038LogicalResult tosa::ArgMaxOp::verify() { return argMaxMinVerify(*this); }
1039
1040LogicalResult tosa::ArgMinOp::verify() { return argMaxMinVerify(*this); }
1041
1042static LogicalResult verifyPoolingOpImpl(Operation *op,
1043 ArrayRef<int64_t> kernel,
1044 ArrayRef<int64_t> strides,
1045 ArrayRef<int64_t> padding, Value input,
1046 Value output) {
1047 if (failed(verifySameElementTypes(op, input.getType(), output.getType())))
1048 return failure();
1049
1050 const bool hasKernel = kernel.size() > 0;
1051 const bool hasStrides = strides.size() > 0;
1052 const bool hasPad = padding.size() > 0;
1053
1054 if (hasKernel && llvm::any_of(kernel, [](int64_t s) { return s < 1; }))
1055 return op->emitOpError("expect all kernel values to be >= 1, got ")
1056 << kernel;
1057
1058 if (hasStrides && llvm::any_of(strides, [](int64_t s) { return s < 1; }))
1059 return op->emitOpError("expect all stride values to be >= 1, got ")
1060 << strides;
1061
1062 if (hasPad && llvm::any_of(padding, [](int64_t p) { return p < 0; }))
1063 return op->emitOpError("expect all padding values to be >= 0, got ")
1064 << padding;
1065
1066 if (hasKernel && hasPad) {
1067 // Padding must be less than kernel size to avoid a divide-by-zero
1068 const int64_t kernelX = kernel[1];
1069 const int64_t padLeft = padding[2];
1070 const int64_t padRight = padding[3];
1071 if (padRight >= kernelX || padLeft >= kernelX)
1072 return op->emitOpError("expected left/right padding to be less than the "
1073 "width of the kernel, got pad_left=")
1074 << padLeft << ", pad_right=" << padRight
1075 << ", kernel_x=" << kernelX;
1076
1077 const int64_t kernelY = kernel[0];
1078 const int64_t padTop = padding[0];
1079 const int64_t padBottom = padding[1];
1080 if (padTop >= kernelY || padBottom >= kernelY)
1081 return op->emitOpError("expected top/bottom padding to be less than the "
1082 "height of the kernel, got pad_top=")
1083 << padTop << ", pad_bottom=" << padBottom
1084 << ", kernel_y=" << kernelY;
1085 }
1086
1087 const auto inputType = llvm::dyn_cast<RankedTensorType>(input.getType());
1088 const auto outputType = llvm::dyn_cast<RankedTensorType>(output.getType());
1089 if (!inputType || !outputType)
1090 return success();
1091
1092 if (hasKernel && hasStrides && hasPad) {
1093 const auto verifyOutputSize =
1094 [op](const int64_t inputSize, const int64_t outputSize,
1095 const int64_t kernelSize, const int64_t strideSize,
1096 const int64_t padBefore, const int64_t padAfter,
1097 const llvm::StringRef dimName, const llvm::StringRef dimAxis,
1098 const llvm::StringRef padBeforeName,
1099 const llvm::StringRef padAfterName) -> LogicalResult {
1100 if (ShapedType::isDynamic(inputSize))
1101 return success();
1102
1103 const std::optional<int64_t> calculatedOutSizeMinusOne =
1104 idivCheck(inputSize + padBefore + padAfter - kernelSize, strideSize);
1105 if (!calculatedOutSizeMinusOne.has_value())
1106 return op->emitOpError("expected input_")
1107 << dimName << " + pad_" << padBeforeName << " + pad_"
1108 << padAfterName << " - kernel_" << dimAxis
1109 << " to be wholly divisible by stride_" << dimAxis << ", got ("
1110 << inputSize << " + " << padBefore << " + " << padAfter << " - "
1111 << kernelSize << ") / " << strideSize;
1112
1113 const int64_t calculatedOutSize = calculatedOutSizeMinusOne.value() + 1;
1114 if (ShapedType::isStatic(outputSize) && calculatedOutSize != outputSize)
1115 return op->emitOpError("calculated output ")
1116 << dimName << " did not match expected: " << "calculated="
1117 << calculatedOutSize << ", expected=" << outputSize;
1118
1119 return success();
1120 };
1121
1122 if (failed(verifyOutputSize(inputType.getDimSize(1),
1123 outputType.getDimSize(1), kernel[0], strides[0],
1124 padding[0], padding[1], "height", "y", "top",
1125 "bottom")))
1126 return failure();
1127
1128 if (failed(verifyOutputSize(
1129 inputType.getDimSize(2), outputType.getDimSize(2), kernel[1],
1130 strides[1], padding[2], padding[3], "width", "x", "left", "right")))
1131 return failure();
1132 }
1133 return success();
1134}
1135
1136template <typename T>
1137static LogicalResult verifyPoolingOp(T op) {
1138 return verifyPoolingOpImpl(op.getOperation(), op.getKernel(), op.getStride(),
1139 op.getPad(), op.getInput(), op.getOutput());
1140}
1141
1142template <typename T>
1143static LogicalResult verifyAvgPoolCommonTypeAndZpChecks(T op) {
1144 const Type inputETy = getStorageElementTypeOrSelf(op.getInput().getType());
1145 const Type resultETy = getStorageElementTypeOrSelf(op.getOutput().getType());
1146 const Type inputZpETy =
1147 getStorageElementTypeOrSelf(op.getInputZp().getType());
1148 const Type outputZpETy =
1149 getStorageElementTypeOrSelf(op.getOutputZp().getType());
1150
1151 auto accType = op.getAccType();
1152 if (llvm::isa<IntegerType>(inputETy) && !accType.isInteger(32))
1153 return op.emitOpError("accumulator type for integer tensor is not i32");
1154
1155 if (inputETy.isF16() && !(accType.isF16() || accType.isF32()))
1156 return op.emitOpError("accumulator type for f16 tensor is not f16/f32");
1157
1158 if (inputETy.isBF16() && !accType.isF32())
1159 return op.emitOpError("accumulator type for bf16 tensor is not f32");
1160
1161 if (inputETy.isF32() && !accType.isF32())
1162 return op.emitOpError("accumulator type for f32 tensor is not f32");
1163
1164 if (inputETy != inputZpETy)
1165 return op.emitOpError("expect both input and its zero point are the same "
1166 "element type, got ")
1167 << inputETy << " and " << inputZpETy;
1168
1169 if (resultETy != outputZpETy)
1170 return op.emitOpError("expect both output and its zero point are the same "
1171 "element type, got ")
1172 << resultETy << " and " << outputZpETy;
1173
1174 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1175 if (succeeded(maybeIZp) && op.verifyInputZeroPoint(*maybeIZp).failed())
1176 return failure();
1177
1178 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1179 if (succeeded(maybeOZp) && op.verifyOutputZeroPoint(*maybeOZp).failed())
1180 return failure();
1181
1182 return success();
1183}
1184
1185namespace {
1186struct AdaptivePoolingConstShapeValues {
1187 llvm::SmallVector<int64_t> kernel;
1188 llvm::SmallVector<int64_t> stride;
1189 llvm::SmallVector<int64_t> pad;
1190};
1191} // namespace
1192
1193template <typename T>
1195 std::is_same_v<T, tosa::AvgPool2dAdaptiveOp> ||
1196 std::is_same_v<T, tosa::MaxPool2dAdaptiveOp>;
1197
1198template <typename T,
1199 typename std::enable_if<IsSupportedAdaptivePoolConstShapeVerifyOp<T>,
1200 int>::type = 0>
1202 T op, AdaptivePoolingConstShapeValues &values) {
1203 tosa::getConstShapeValues(op.getKernel().getDefiningOp(), values.kernel);
1204 tosa::getConstShapeValues(op.getStride().getDefiningOp(), values.stride);
1205 tosa::getConstShapeValues(op.getPad().getDefiningOp(), values.pad);
1206}
1207
1208LogicalResult tosa::AvgPool2dOp::verify() {
1209 if (failed(verifyPoolingOp(*this)))
1210 return failure();
1212 return failure();
1213 return success();
1214}
1215
1216LogicalResult tosa::AvgPool2dAdaptiveOp::verify() {
1217 AdaptivePoolingConstShapeValues values;
1219
1220 // If pad/stride/kernel are not constant, this is okay, we just can't check
1221 // their values. extractAdaptivePoolingConstShapeOperands will return an empty
1222 // list for each non CTC input. verifyPoolingOpImpl will need to handle values
1223 // not being present, and return success if they cannot be checked.
1224
1225 if (failed(verifyPoolingOpImpl(getOperation(), values.kernel, values.stride,
1226 values.pad, getInput(), getOutput())))
1227 return failure();
1228
1230 return failure();
1231
1232 return success();
1233}
1234
1235LogicalResult tosa::ClampOp::verify() {
1236 mlir::Type inputETy =
1237 llvm::cast<ShapedType>(getInput().getType()).getElementType();
1238 if (auto quantType =
1239 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(inputETy)) {
1240 inputETy = getStorageElementTypeFromQuantized(quantType);
1241 }
1242 mlir::Type outputETy =
1243 llvm::cast<ShapedType>(getOutput().getType()).getElementType();
1244 if (auto quantType =
1245 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(outputETy)) {
1246 outputETy = getStorageElementTypeFromQuantized(quantType);
1247 }
1248 if (inputETy != outputETy)
1249 return emitOpError("input/output element types are incompatible.");
1250
1251 auto maxValAttr = getMaxValAttr();
1252 auto minValAttr = getMinValAttr();
1253
1254 unsigned dataTypeBitWidth = inputETy.getIntOrFloatBitWidth();
1255
1256 if (inputETy.isInteger(dataTypeBitWidth)) {
1257 // if input datatype is integer, check that the min_val/max_val attributes
1258 // are integer attributes, and that their type is the same as the input's
1259 // datatype
1260 auto intMaxValAttr = mlir::dyn_cast<mlir::IntegerAttr>(maxValAttr);
1261 auto intMinValAttr = mlir::dyn_cast<mlir::IntegerAttr>(minValAttr);
1262 if (!intMaxValAttr || !intMinValAttr ||
1263 (intMaxValAttr.getType() != intMinValAttr.getType()) ||
1264 (intMaxValAttr.getType() != inputETy))
1265 return emitOpError("min/max attributes types are incompatible with "
1266 "input/output element types.");
1267
1268 const bool isUnsigned = inputETy.isUnsignedInteger();
1269 const bool isBoolean = inputETy.isInteger(1);
1270 const APInt minVal = intMinValAttr.getValue();
1271 const APInt maxVal = intMaxValAttr.getValue();
1272 if ((isUnsigned || isBoolean) ? maxVal.ult(minVal) : maxVal.slt(minVal))
1273 return emitOpError("expected min_val <= max_val, got min_val=")
1274 << minValAttr << ", max_val=" << maxValAttr;
1275 } else {
1276 // otherwise, input datatype is float, check that the min_val/max_val
1277 // attributes share the same type and that their type is the same as the
1278 // input's datatype
1279 auto floatMaxValAttr = mlir::dyn_cast<mlir::FloatAttr>(maxValAttr);
1280 auto floatMinValAttr = mlir::dyn_cast<mlir::FloatAttr>(minValAttr);
1281 if (!floatMaxValAttr || !floatMinValAttr ||
1282 (floatMaxValAttr.getType() != floatMinValAttr.getType()) ||
1283 (floatMaxValAttr.getType() != inputETy))
1284 return emitOpError("min/max attributes types are incompatible with "
1285 "input/output element types.");
1286
1287 const APFloat minVal = floatMinValAttr.getValue();
1288 const APFloat maxVal = floatMaxValAttr.getValue();
1289 if (minVal.isNaN() || maxVal.isNaN())
1290 return emitOpError("min/max attributes should not be 'NaN', got min_val=")
1291 << minValAttr << ", max_val=" << maxValAttr;
1292
1293 if (maxVal < minVal)
1294 return emitOpError("expected min_val <= max_val, got min_val=")
1295 << minValAttr << ", max_val=" << maxValAttr;
1296 }
1297
1298 return success();
1299}
1300
1301//===----------------------------------------------------------------------===//
1302// TOSA Operator Quantization Builders.
1303//===----------------------------------------------------------------------===//
1304
1305/// This builder is called on all convolution operators except TransposeConv,
1306/// which has specialized output shape semantics. The builder also defines the
1307/// bitwidth of the output given the bit width of the input & weight content.
1309 Type outputType, Value input, Value weight,
1310 Value bias, DenseI64ArrayAttr pad,
1311 DenseI64ArrayAttr stride,
1312 DenseI64ArrayAttr dilation,
1313 TypeAttr accType) {
1314 auto zps = createZPsAsConst(builder, input, weight);
1315 result.addOperands({input, weight, bias, zps.first, zps.second});
1316 result.addAttribute("pad", pad);
1317 result.addAttribute("stride", stride);
1318 result.addAttribute("dilation", dilation);
1319 result.addAttribute("acc_type", accType);
1320 Type finalOutputType = outputType;
1321 auto quantAttr = buildConvOpQuantizationAttr(builder, input, weight);
1322 if (quantAttr) {
1323 finalOutputType =
1324 buildConvOpResultTypeInfo(builder, outputType, input, weight);
1325 }
1326 result.addTypes(finalOutputType);
1327}
1328
1329/// Handles tosa.transpose_conv2d which has outpad and output shape
1330/// attributes.
1331static void
1333 Type outputType, Value input, Value weight,
1334 Value bias, DenseI64ArrayAttr outpad,
1335 DenseI64ArrayAttr stride, TypeAttr accType) {
1336 auto zps = createZPsAsConst(builder, input, weight);
1337 result.addOperands({input, weight, bias, zps.first, zps.second});
1338 result.addAttribute("out_pad", outpad);
1339 result.addAttribute("stride", stride);
1340 result.addAttribute("acc_type", accType);
1341 Type finalOutputType = outputType;
1342 auto quantAttr = buildConvOpQuantizationAttr(builder, input, weight);
1343 if (quantAttr) {
1344 finalOutputType =
1345 buildConvOpResultTypeInfo(builder, outputType, input, weight);
1346 }
1347 result.addTypes(finalOutputType);
1348}
1349
1352 Type outputType, Value a, Value b) {
1353 const std::pair<Value, Value> zps = createZPsAsConst(builder, a, b);
1354 result.addOperands({a, b, zps.first, zps.second});
1355
1356 Type finalOutputType{outputType};
1357 if (buildMatMulOpQuantizationAttr(builder, a, b)) {
1358 auto eType = getStorageElementTypeOrSelf(a.getType());
1359 auto inputBits = eType.getIntOrFloatBitWidth();
1360
1361 auto outputShapedType = llvm::dyn_cast<ShapedType>(outputType);
1362 assert(outputShapedType && "Output must be a shaped type");
1363
1364 IntegerType accElementType;
1365 if (inputBits == 16)
1366 accElementType = builder.getIntegerType(48);
1367 else
1368 accElementType = builder.getI32Type();
1369
1370 finalOutputType = outputShapedType.clone(accElementType);
1371 }
1372 result.addTypes(finalOutputType);
1373}
1374
1376 OperationState &result, Type outputType,
1377 Value a, Value b) {
1378 buildMatMulLikeOpWithQuantInfo(builder, result, outputType, a, b);
1379}
1380
1382 OperationState &result, Type outputType,
1383 Value a, Value b) {
1384 buildMatMulLikeOpWithQuantInfo(builder, result, outputType, a, b);
1385}
1386
1387/// Both the tosa.avg_pool2d and unary ops use the same
1388/// UnaryOpQuantizationAttr but avg_pool operator has its own builder as it
1389/// has additional parameters not part of the unary ops.
1390static void
1392 Type outputType, Value input,
1393 DenseArrayAttr kernel, DenseArrayAttr stride,
1394 DenseArrayAttr pad, TypeAttr accType) {
1395 const Location loc{result.location};
1396 int64_t inputZp{0};
1397 int64_t outputZp{0};
1398
1399 if (auto quantAttr =
1400 buildUnaryOpQuantizationAttr(builder, input, outputType)) {
1401 inputZp = quantAttr.getInputZp();
1402 outputZp = quantAttr.getOutputZp();
1403 }
1404 const std::optional<Value> inputZpOp =
1405 createZeroPointTensor(builder, loc, input.getType(), inputZp);
1406 if (!inputZpOp) {
1407 (void)emitError(
1408 loc,
1409 "Failed to create input zero point tensor for quantized AVG_POOL2D op");
1410 }
1411 const std::optional<Value> outputZpOp =
1412 createZeroPointTensor(builder, loc, outputType, outputZp);
1413 if (!outputZpOp) {
1414 (void)emitError(loc, "Failed to create output zero point tensor for "
1415 "quantized AVG_POOL2D op");
1416 }
1417
1418 if (inputZpOp && outputZpOp) {
1419 result.addOperands({input, inputZpOp.value(), outputZpOp.value()});
1420 } else {
1421 // failed to create one or more zero points above: just add input as
1422 // operands this will trigger error in building the op because of missing
1423 // zero points
1424 result.addOperands({input});
1425 }
1426 result.addAttribute("kernel", kernel);
1427 result.addAttribute("stride", stride);
1428 result.addAttribute("pad", pad);
1429 result.addAttribute("acc_type", accType);
1430 result.types.push_back(outputType);
1431}
1432
1433/// This builder mirrors avg_pool2d quant-info handling and materializes
1434/// kernel/stride/pad as const_shape operands for avg_pool2d_adaptive.
1436 OpBuilder &builder, OperationState &result, Type outputType, Value input,
1438 TypeAttr accType) {
1439 const Location loc{result.location};
1440 int64_t inputZp{0};
1441 int64_t outputZp{0};
1442
1443 if (auto quantAttr =
1444 buildUnaryOpQuantizationAttr(builder, input, outputType)) {
1445 inputZp = quantAttr.getInputZp();
1446 outputZp = quantAttr.getOutputZp();
1447 }
1448 const std::optional<Value> inputZpOp =
1449 createZeroPointTensor(builder, loc, input.getType(), inputZp);
1450 if (!inputZpOp) {
1451 (void)emitError(loc,
1452 "Failed to create input zero point tensor for quantized "
1453 "AVG_POOL2D_ADAPTIVE op");
1454 }
1455 const std::optional<Value> outputZpOp =
1456 createZeroPointTensor(builder, loc, outputType, outputZp);
1457 if (!outputZpOp) {
1458 (void)emitError(loc, "Failed to create output zero point tensor for "
1459 "quantized AVG_POOL2D_ADAPTIVE op");
1460 }
1461
1462 if (inputZpOp && outputZpOp) {
1463 ImplicitLocOpBuilder b(loc, builder);
1464 Value kernelShape = getTosaConstShape(b, kernel.asArrayRef());
1465 Value strideShape = getTosaConstShape(b, stride.asArrayRef());
1466 Value padShape = getTosaConstShape(b, pad.asArrayRef());
1467 result.addOperands({input, inputZpOp.value(), outputZpOp.value(),
1468 kernelShape, strideShape, padShape});
1469 } else {
1470 // Failed to create one or more zero points above: just add input as
1471 // operands. This will trigger error in building the op because of missing
1472 // operands.
1473 result.addOperands({input});
1474 }
1475 result.addAttribute("acc_type", accType);
1476 result.types.push_back(outputType);
1477}
1478
1479/// This builder is called on single-parameter negate operator
1480/// to construct input and output zero points based on their
1481/// types.
1483 OperationState &result, Type outputType,
1484 Value input) {
1485 const Location loc{result.location};
1486 int64_t input1Zp{0};
1487 int64_t outputZp{0};
1488 auto quantAttr = buildUnaryOpQuantizationAttr(builder, input, outputType);
1489 if (quantAttr) {
1490 input1Zp = quantAttr.getInputZp();
1491 outputZp = quantAttr.getOutputZp();
1492 }
1493 const std::optional<Value> input1ZpOp =
1494 createZeroPointTensor(builder, loc, input.getType(), input1Zp);
1495 if (!input1ZpOp) {
1496 (void)emitError(
1497 loc, "Failed to create input1 zero point for quantized NEGATE op");
1498 }
1499
1500 const std::optional<Value> outputZpOp =
1501 createZeroPointTensor(builder, loc, input.getType(), outputZp);
1502 if (!outputZpOp) {
1503 (void)emitError(
1504 loc, "Failed to create output zero point for quantized NEGATE op");
1505 }
1506
1507 if (input1ZpOp && outputZpOp) {
1508 result.addOperands({input, input1ZpOp.value(), outputZpOp.value()});
1509 } else {
1510 // failed to create one or more zero points above: just add input as
1511 // operands. This will trigger error in building the op because of
1512 // missing zero points
1513 result.addOperands({input});
1514 }
1515
1516 result.types.push_back(outputType);
1517}
1518
1519/// This builder is called on TOSA pad operator that needs to create its own
1520/// OptionalAttr quantization_attr parameter to scale the padding values
1521/// correctly. No pad_const is interpreted as zero-padding.
1523 Type outputType, Value input,
1524 Value paddings) {
1525 const Location loc{result.location};
1526 int32_t zp{0};
1527 const auto quantAttr = buildPadOpQuantizationAttr(builder, input);
1528 if (quantAttr) {
1529 zp = static_cast<int32_t>(quantAttr.getInputZp());
1530 }
1531 const auto padConstOp{createPadConstTensor(builder, loc, input, zp)};
1532 result.addOperands({input, paddings, padConstOp});
1533 result.types.push_back(outputType);
1534}
1535
1537 StringRef name, Type variableType,
1538 Attribute initialValue) {
1539 const Location loc{result.location};
1540 auto nameAttr = builder.getStringAttr(name);
1541
1542 auto shapedType = dyn_cast<ShapedType>(variableType);
1543 if (!shapedType) {
1544 (void)emitError(loc, "variable type must be a shaped type");
1545 return;
1546 }
1547 if (!shapedType.hasRank()) {
1548 (void)emitError(loc, "variable type must be a ranked type");
1549 return;
1550 }
1551
1552 auto elementType = shapedType.getElementType();
1553 auto elementTypeAttr = TypeAttr::get(elementType);
1554 ArrayRef<int64_t> shape = shapedType.getShape();
1555 auto varShapeAttr = builder.getIndexTensorAttr(convertFromMlirShape(shape));
1556
1557 result.addAttribute("sym_name", nameAttr);
1558 result.addAttribute("var_shape", varShapeAttr);
1559 result.addAttribute("type", elementTypeAttr);
1560 result.addAttribute("initial_value", initialValue);
1561}
1562
1563//===----------------------------------------------------------------------===//
1564// TOSA Operator Return Type Inference.
1565//===----------------------------------------------------------------------===//
1566static FailureOr<int64_t> resolveBroadcastDim(const int64_t dim1,
1567 const int64_t dim2) {
1568 if (dim1 == 1)
1569 return dim2;
1570 if (dim2 == 1)
1571 return dim1;
1572
1573 if (ShapedType::isStatic(dim1) && ShapedType::isStatic(dim2) && dim1 != dim2)
1574 return failure();
1575
1576 // Prefer static dimension over dynamic
1577 return ShapedType::isDynamic(dim1) ? dim2 : dim1;
1578}
1579
1580static LogicalResult resolveBroadcastShape(const ValueShapeRange &operands,
1581 SmallVector<int64_t> &outShape) {
1582 int64_t outRank = 0;
1583 for (int i = 0, e = operands.size(); i != e; ++i) {
1584 auto shape = operands.getShape(i);
1585 if (!shape.hasRank()) {
1586 // TODO(jennik): Update function to have better case handling for
1587 // invalid operands and for ranked tensors.
1588 return failure();
1589 }
1590 outRank = std::max<int64_t>(outRank, shape.getRank());
1591 }
1592
1593 outShape.resize(outRank, 1);
1594
1595 for (int i = 0, e = operands.size(); i != e; ++i) {
1596 auto shape = operands.getShape(i);
1597 auto rankDiff = outShape.size() - shape.getRank();
1598
1599 for (size_t i = 0, e = shape.getRank(); i < e; ++i) {
1600 auto dim1 = outShape[i + rankDiff];
1601 auto dim2 = shape.getDimSize(i);
1602
1603 const FailureOr<int64_t> maybeResolvedDim =
1604 resolveBroadcastDim(dim1, dim2);
1605 if (failed(maybeResolvedDim))
1606 return failure();
1607 const int64_t resolvedDim = *maybeResolvedDim;
1608 outShape[i + rankDiff] = resolvedDim;
1609 }
1610 }
1611
1612 return success();
1613}
1614
1615LogicalResult tosa::ArgMaxOp::inferReturnTypeComponents(
1616 MLIRContext *context, ::std::optional<Location> location,
1617 ArgMaxOp::Adaptor adaptor,
1618 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
1619 return inferArgMaxMinReturnTypeComponents(context, location, adaptor,
1620 inferredReturnShapes);
1621}
1622
1623LogicalResult tosa::ArgMinOp::inferReturnTypeComponents(
1624 MLIRContext *context, ::std::optional<Location> location,
1625 ArgMinOp::Adaptor adaptor,
1626 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
1627 return inferArgMaxMinReturnTypeComponents(context, location, adaptor,
1628 inferredReturnShapes);
1629}
1630
1631LogicalResult tosa::RFFT2dOp::inferReturnTypeComponents(
1632 MLIRContext *context, ::std::optional<Location> location,
1633 RFFT2dOp::Adaptor adaptor,
1634 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
1635 ShapeAdaptor inputShape(adaptor.getInputReal().getType());
1636
1637 if (!inputShape.hasRank())
1638 return failure();
1639
1640 llvm::SmallVector<int64_t> outputShape;
1641 outputShape.resize(3, ShapedType::kDynamic);
1642 outputShape[0] = inputShape.getDimSize(0);
1643 outputShape[1] = inputShape.getDimSize(1);
1644 int64_t inWidth = inputShape.getDimSize(2);
1645
1646 // Note that we can support this calculation symbolically
1647 // in the future e.g. [x, y, z] -> [x, y, z / 2 + 1]
1648 if (inWidth != ShapedType::kDynamic)
1649 outputShape[2] = inWidth / 2 + 1;
1650
1651 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
1652 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
1653
1654 return success();
1655}
1656
1657static LogicalResult verifyDimIsPowerOfTwo(Operation *op, const int64_t dimSize,
1658 const llvm::StringRef dimName) {
1659 const bool isPowerOfTwo = (dimSize & (dimSize - 1)) == 0 && dimSize > 0;
1660 if (!isPowerOfTwo)
1661 return op->emitOpError("expected ")
1662 << dimName << " to be a power of two, got " << dimSize;
1663
1664 return success();
1665}
1666
1667LogicalResult tosa::RFFT2dOp::verify() {
1668 const auto outputTypes = getResultTypes();
1669 if (failed(verifyCompatibleShapes(outputTypes)))
1670 return emitOpError("expected output shapes to match, got ") << outputTypes;
1671
1672 const auto inputType =
1673 llvm::dyn_cast<RankedTensorType>(getInputReal().getType());
1674 if (!inputType)
1675 return success();
1676
1677 const int64_t height = inputType.getDimSize(1);
1678 if (ShapedType::isStatic(height) &&
1679 failed(verifyDimIsPowerOfTwo(*this, height, "height")))
1680 return failure();
1681
1682 const int64_t width = inputType.getDimSize(2);
1683 if (ShapedType::isStatic(width) &&
1684 failed(verifyDimIsPowerOfTwo(*this, width, "width")))
1685 return failure();
1686
1687 const auto outputType = llvm::dyn_cast<RankedTensorType>(outputTypes[0]);
1688 if (!outputType)
1689 return success();
1690
1691 // Batch and height input/output dimensions should match
1692 if (failed(verifyCompatibleShape(inputType.getShape().drop_back(),
1693 outputType.getShape().drop_back())))
1694 return emitOpError("expected batch and height dimensions of input/output "
1695 "to match, got input=")
1696 << inputType << " output=" << outputType;
1697
1698 // Output width dimension expected to be input_width / 2 + 1
1699 const int64_t outputWidth = outputType.getDimSize(2);
1700 if (ShapedType::isStatic(width) && ShapedType::isStatic(outputWidth) &&
1701 (outputWidth != (width / 2) + 1))
1702 return emitOpError(
1703 "expected output width to be equal to input_width / 2 + 1, got ")
1704 << outputWidth;
1705
1706 return success();
1707}
1708
1709LogicalResult tosa::FFT2dOp::inferReturnTypeComponents(
1710 MLIRContext *context, ::std::optional<Location> location,
1711 FFT2dOp::Adaptor adaptor,
1712 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
1713 inferredReturnShapes.push_back(
1714 ShapedTypeComponents(ShapeAdaptor(adaptor.getInputReal().getType())));
1715 inferredReturnShapes.push_back(
1716 ShapedTypeComponents(ShapeAdaptor(adaptor.getInputImag().getType())));
1717 return success();
1718}
1719
1720LogicalResult tosa::FFT2dOp::verify() {
1721 const auto inputRealType =
1722 llvm::dyn_cast<RankedTensorType>(getInputReal().getType());
1723 const auto inputImagType =
1724 llvm::dyn_cast<RankedTensorType>(getInputImag().getType());
1725 if (!inputRealType || !inputImagType)
1726 return success();
1727
1728 const auto trySelectStaticDim = [](const int64_t a, const int64_t b) {
1729 return ShapedType::isDynamic(a) ? a : b;
1730 };
1731
1732 const int64_t height = trySelectStaticDim(inputRealType.getDimSize(1),
1733 inputImagType.getDimSize(1));
1734 if (ShapedType::isStatic(height) &&
1735 failed(verifyDimIsPowerOfTwo(*this, height, "height")))
1736 return failure();
1737
1738 const int64_t width = trySelectStaticDim(inputRealType.getDimSize(2),
1739 inputImagType.getDimSize(2));
1740 if (ShapedType::isStatic(width) &&
1741 failed(verifyDimIsPowerOfTwo(*this, width, "width")))
1742 return failure();
1743
1744 return success();
1745}
1746
1747LogicalResult tosa::ConcatOp::inferReturnTypeComponents(
1748 MLIRContext *context, ::std::optional<Location> location,
1749 ConcatOp::Adaptor adaptor,
1750 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
1751 // Infer all dimension sizes by reducing based on inputs.
1752 const Properties &prop = adaptor.getProperties();
1753 int32_t axis = prop.axis.getValue().getSExtValue();
1754 llvm::SmallVector<int64_t> outputShape;
1755 bool hasRankedInput = false;
1756 for (auto operand : adaptor.getOperands()) {
1757 ShapeAdaptor operandShape(operand.getType());
1758 if (!operandShape.hasRank())
1759 continue;
1760
1761 // Copy the Operand's rank.
1762 if (!hasRankedInput)
1763 outputShape.resize(operandShape.getRank(), ShapedType::kDynamic);
1764
1765 // Copy shapes until the dim is non-dynamic.
1766 for (int i = 0, s = operandShape.getRank(); i < s; i++) {
1767 if (i == axis || operandShape.isDynamicDim(i))
1768 continue;
1769 if (outputShape[i] == ShapedType::kDynamic)
1770 outputShape[i] = operandShape.getDimSize(i);
1771 if (outputShape[i] != operandShape.getDimSize(i))
1772 return emitOptionalError(location,
1773 "Cannot concat tensors with different sizes"
1774 " on the non-axis dimension ",
1775 i);
1776 }
1777
1778 hasRankedInput = true;
1779 }
1780
1781 if (adaptor.getInput1().empty())
1782 return failure();
1783
1784 Type inputType =
1785 llvm::cast<TensorType>(adaptor.getInput1().getType()[0]).getElementType();
1786 if (!hasRankedInput) {
1787 inferredReturnShapes.push_back(ShapedTypeComponents(inputType));
1788 return success();
1789 }
1790
1791 // Determine the dimension size along the concatenation axis.
1792 int64_t concatDimSize = 0;
1793 for (auto operand : adaptor.getOperands()) {
1794 ShapeAdaptor operandShape(operand.getType());
1795
1796 // We need to know the length of the concatenation axis of all inputs to
1797 // determine the dimension size of the output shape.
1798 if (!operandShape.hasRank() || operandShape.isDynamicDim(axis)) {
1799 concatDimSize = ShapedType::kDynamic;
1800 break;
1801 }
1802
1803 concatDimSize += operandShape.getDimSize(axis);
1804 }
1805
1806 outputShape[axis] = concatDimSize;
1807
1808 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape, inputType));
1809 return success();
1810}
1811
1812LogicalResult tosa::ConcatOp::verify() {
1813 // check that each input has same element type as output
1814 auto outType = getOutput().getType();
1815 const Operation::operand_range inputList = getInput1();
1816
1817 // Check there is at least one input
1818 if (inputList.empty())
1819 return emitOpError("expect at least one input");
1820
1821 if (!llvm::all_of(inputList, [&](auto input) {
1822 return succeeded(verifySameElementTypes(
1823 *this, /* inType = */ input.getType(), outType));
1824 })) {
1825 return failure();
1826 }
1827
1828 const int32_t axis = getAxis();
1829 ShapeAdaptor firstRankedInputShape = nullptr;
1830 for (const auto &input : inputList) {
1831 const Type inputType = input.getType();
1832 ShapeAdaptor currShape(inputType);
1833 if (currShape.hasRank()) {
1834 firstRankedInputShape = currShape;
1835 // Check axis is in expected range
1836 if (axis < 0 || axis >= firstRankedInputShape.getRank())
1837 return emitOpError("expect axis to be within range 0 < axis < "
1838 "rank(input1[firstRankedTensorIdx]), got ")
1839 << axis;
1840 break;
1841 }
1842 }
1843
1844 const auto allOperandsHasRank = [](const Value input) {
1845 return ShapeAdaptor(input.getType()).hasRank();
1846 };
1847 if (llvm::all_of(inputList, allOperandsHasRank)) {
1848 const int64_t firstInputRank = firstRankedInputShape.getRank();
1849
1850 for (const auto &[index, input] : llvm::enumerate(inputList.drop_front())) {
1851 const ShapeAdaptor inputShape(input.getType());
1852 const int64_t inputRank = inputShape.getRank();
1853 const size_t operandNum = index + 1;
1854
1855 // Check that each operand has the same rank
1856 if (inputRank != firstInputRank)
1857 return emitOpError(
1858 "expect all operands to have the same rank, but got ")
1859 << firstInputRank << " vs " << inputRank << " on operands 0 and "
1860 << operandNum;
1861
1862 // Check non-axis dims match
1863 for (int i = 0; i < inputRank; i++) {
1864 const int64_t inputDim = inputShape.getDimSize(i);
1865 const int64_t firstInputDim = firstRankedInputShape.getDimSize(i);
1866 if (i == axis || firstRankedInputShape.isDynamicDim(i) ||
1867 inputShape.isDynamicDim(i))
1868 continue;
1869 if (inputDim != firstInputDim)
1870 return emitOpError("expect all operand shapes to have the same sizes "
1871 "on non-axis dimensions, but got ")
1872 << inputDim << " vs " << firstInputDim << " at index " << i
1873 << " on operands 0 and " << operandNum;
1874 }
1875 }
1876
1877 const ShapeAdaptor outputShape(outType);
1878 if (outputShape.hasRank() && outputShape.getRank() != firstInputRank)
1879 return emitOpError("expect output rank to match inputs rank, got ")
1880 << outputShape.getRank() << " vs " << firstInputRank;
1881
1882 // ERROR_IF(axis_sum != shape[axis]);
1883 int64_t axisSum = 0;
1884 for (const auto &input : inputList) {
1885 const ShapeAdaptor inputShape(input.getType());
1886 if (inputShape.isDynamicDim(axis)) {
1887 // make axisSum negative to indicate invalid value
1888 axisSum = -1;
1889 break;
1890 }
1891 axisSum += inputShape.getDimSize(axis);
1892 }
1893
1894 if (axisSum >= 0 && outputShape.hasRank() &&
1895 !outputShape.isDynamicDim(axis) &&
1896 axisSum != outputShape.getDimSize(axis))
1897 return emitOpError("requires sum of axis dimensions of input1 "
1898 "equal to output axis dimension, got ")
1899 << axisSum << " and " << outputShape.getDimSize(axis);
1900 }
1901
1902 return success();
1903}
1904
1905LogicalResult tosa::EqualOp::inferReturnTypeComponents(
1906 MLIRContext *context, ::std::optional<Location> location,
1907 ValueShapeRange operands, DictionaryAttr attributes, PropertyRef properties,
1908 RegionRange regions,
1909 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
1910 auto elementType = IntegerType::get(context, /*width=*/1);
1911
1913 if (resolveBroadcastShape(operands, outShape).failed()) {
1914 inferredReturnShapes.push_back(ShapedTypeComponents(elementType));
1915 return success();
1916 }
1917
1918 inferredReturnShapes.push_back(ShapedTypeComponents(outShape, elementType));
1919 return success();
1920}
1921
1922bool tosa::EqualOp::isCompatibleReturnTypes(TypeRange l, TypeRange r) {
1923 if (l.size() != r.size() || l.size() != 1)
1924 return false;
1925 return succeeded(verifyCompatibleShape(l[0], r[0]));
1926}
1927
1928// MATMUL batch shapes are right aligned and may have different ranks. Missing
1929// leading dimensions are treated as one, while an unranked input remains
1930// unknown because it may contain any number of batch dimensions.
1932 int64_t axis) {
1933 if (!shape.hasRank())
1934 return ShapedType::kDynamic;
1935 const int64_t inputAxis = axis - (outputRank - shape.getRank());
1936 return inputAxis < 0 ? 1 : shape.getDimSize(inputAxis);
1937}
1938
1939static FailureOr<SmallVector<int64_t>>
1941 int64_t outputRank, bool transposeB) {
1942 if (outputRank < 2 ||
1943 (aShape.hasRank() &&
1944 (aShape.getRank() < 2 || aShape.getRank() > outputRank)) ||
1945 (bShape.hasRank() &&
1946 (bShape.getRank() < 2 || bShape.getRank() > outputRank)))
1947 return failure();
1948
1949 SmallVector<int64_t> outputShape(outputRank, ShapedType::kDynamic);
1950 for (int64_t axis = 0; axis < outputRank - 2; ++axis) {
1951 const int64_t aDim = getMatMulBatchDim(aShape, outputRank, axis);
1952 const int64_t bDim = getMatMulBatchDim(bShape, outputRank, axis);
1953 FailureOr<int64_t> resolvedDim = resolveBroadcastDim(aDim, bDim);
1954 if (failed(resolvedDim))
1955 return failure();
1956 outputShape[axis] = *resolvedDim;
1957 }
1958
1959 if (aShape.hasRank())
1960 outputShape[outputRank - 2] = aShape.getDimSize(aShape.getRank() - 2);
1961 if (bShape.hasRank())
1962 outputShape[outputRank - 1] =
1963 bShape.getDimSize(bShape.getRank() - (transposeB ? 2 : 1));
1964 return outputShape;
1965}
1966
1968 const ShapeAdaptor &aShape, const ShapeAdaptor &bShape, bool transposeB,
1969 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
1970 if ((aShape.hasRank() && aShape.getRank() < 2) ||
1971 (bShape.hasRank() && bShape.getRank() < 2))
1972 return failure();
1973
1974 // An unranked input can have an arbitrary batch prefix, so the output rank
1975 // cannot yet be inferred.
1976 if (!aShape.hasRank() || !bShape.hasRank()) {
1977 inferredReturnShapes.emplace_back();
1978 return success();
1979 }
1980
1981 const int64_t aChannels = aShape.getDimSize(aShape.getRank() - 1);
1982 const int64_t bChannels =
1983 bShape.getDimSize(bShape.getRank() - (transposeB ? 1 : 2));
1984 if (ShapedType::isStatic(aChannels) && ShapedType::isStatic(bChannels) &&
1985 aChannels != bChannels)
1986 return failure();
1987
1988 const int64_t outputRank = std::max(aShape.getRank(), bShape.getRank());
1989 FailureOr<SmallVector<int64_t>> outputShape =
1990 resolveMatMulOutputShape(aShape, bShape, outputRank, transposeB);
1991 if (failed(outputShape))
1992 return failure();
1993
1994 inferredReturnShapes.emplace_back(*outputShape);
1995 return success();
1996}
1997
1998LogicalResult tosa::MatMulOp::inferReturnTypeComponents(
1999 MLIRContext *context, ::std::optional<Location> location,
2000 MatMulOp::Adaptor adaptor,
2001 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2002 return inferMatMulReturnTypeComponents(ShapeAdaptor(adaptor.getA().getType()),
2003 ShapeAdaptor(adaptor.getB().getType()),
2004 /*transposeB=*/false,
2005 inferredReturnShapes);
2006}
2007
2008template <typename T>
2009static LogicalResult verifyMatMulQuantizedOperandsType(T op, Type aElementType,
2010 Type bElementType) {
2011 const auto aQuantizedEType =
2012 llvm::dyn_cast<quant::UniformQuantizedType>(aElementType);
2013 const auto bQuantizedEType =
2014 llvm::dyn_cast<quant::UniformQuantizedType>(bElementType);
2015
2016 if (aQuantizedEType || bQuantizedEType) {
2017 if (!aQuantizedEType || !bQuantizedEType) {
2018 return op.emitOpError("expect operands to be both quantized or both not "
2019 "quantized, got ")
2020 << aElementType << " and " << bElementType;
2021 }
2022 // both a and b have quantized element types
2023 auto aQuantWidth = aQuantizedEType.getStorageTypeIntegralWidth();
2024 auto bQuantWidth = bQuantizedEType.getStorageTypeIntegralWidth();
2025 if (aQuantWidth != bQuantWidth) {
2026 return op.emitOpError("expect quantized operands to have same widths, "
2027 "got ")
2028 << aQuantWidth << " and " << bQuantWidth;
2029 }
2030 }
2031
2032 return success();
2033}
2034
2035template <typename T>
2036static LogicalResult verifyMatMulZeroPointType(T op, Value input, Value zp,
2037 StringRef inputName,
2038 StringRef zpName) {
2039 const Type inputElementType = getElementTypeOrSelf(input.getType());
2040 const Type inputStorageElementType = getStorageElementTypeOrSelf(input);
2041 const Type zpElementType = getStorageElementTypeOrSelf(zp);
2042 Type expectedElementType = inputStorageElementType;
2043
2044 if (isa<BlockScaledType>(inputElementType))
2045 expectedElementType = Float32Type::get(op.getContext());
2046
2047 if (expectedElementType == zpElementType)
2048 return success();
2049
2050 InFlightDiagnostic diag = op.emitOpError("expect input ");
2051 diag << inputName << " and " << zpName;
2052 if (isa<BlockScaledType>(inputElementType))
2053 diag << " have compatible element types, got " << inputElementType
2054 << " and " << zpElementType;
2055 else
2056 diag << " have the same element type, got " << inputStorageElementType
2057 << " and " << zpElementType;
2058 return diag;
2059}
2060
2062 SmallVector<int64_t> batchShape;
2063 if (!shape.hasRank() || shape.getRank() < 2)
2064 return batchShape;
2065 batchShape.reserve(shape.getRank() - 2);
2066 for (int64_t i = 0, e = shape.getRank() - 2; i < e; ++i)
2067 batchShape.push_back(shape.getDimSize(i));
2068 return batchShape;
2069}
2070
2071template <typename T>
2072static LogicalResult verifyMatMulShapes(T op, bool transposeB) {
2073 const ShapeAdaptor aShape(op.getA().getType());
2074 const ShapeAdaptor bShape(op.getB().getType());
2075 const auto outputType = cast<ShapedType>(op.getResult().getType());
2076
2077 int64_t channels = aShape.hasRank() ? aShape.getDimSize(aShape.getRank() - 1)
2078 : ShapedType::kDynamic;
2079 if (bShape.hasRank() &&
2080 failed(tryUpdateDimOrFailure(
2081 op, channels,
2082 bShape.getDimSize(bShape.getRank() - (transposeB ? 1 : 2)), "b",
2083 "channels")))
2084 return failure();
2085
2086 const int64_t minimumOutputRank =
2087 std::max(aShape.hasRank() ? aShape.getRank() : 2,
2088 bShape.hasRank() ? bShape.getRank() : 2);
2089 if (outputType.hasRank() && outputType.getRank() < minimumOutputRank)
2090 return op.emitOpError("expected output rank of at least ")
2091 << minimumOutputRank << ", got " << outputType.getRank();
2092
2093 const bool bothInputsRanked = aShape.hasRank() && bShape.hasRank();
2094 if (!bothInputsRanked && !outputType.hasRank())
2095 return success();
2096
2097 // When an input is unranked, use the declared output rank to validate every
2098 // result dimension constrained by the ranked input.
2099 const int64_t expectedOutputRank =
2100 bothInputsRanked ? minimumOutputRank : outputType.getRank();
2101 FailureOr<SmallVector<int64_t>> expectedOutputShape =
2102 resolveMatMulOutputShape(aShape, bShape, expectedOutputRank, transposeB);
2103 if (failed(expectedOutputShape)) {
2104 InFlightDiagnostic diag = op.emitOpError(
2105 "expected batch dimensions of a and b to be broadcast compatible, "
2106 "got a=[");
2108 diag << "] and b=[";
2110 diag << "]";
2111 return diag;
2112 }
2113
2114 if (outputType.hasRank())
2116 op.getOperation(), outputType, *expectedOutputShape);
2117 return success();
2118}
2119
2120LogicalResult MatMulOp::verify() {
2121 const Type aElementType = getElementTypeOrSelf(getA());
2122 const Type bElementType = getElementTypeOrSelf(getB());
2123
2124 if (failed(
2125 verifyMatMulQuantizedOperandsType(*this, aElementType, bElementType)))
2126 return failure();
2127
2128 if (failed(verifyMatMulZeroPointType(*this, getA(), getAZp(), "a", "a_zp")) ||
2129 failed(verifyMatMulZeroPointType(*this, getB(), getBZp(), "b", "b_zp")))
2130 return failure();
2131
2132 FailureOr<int64_t> maybeAZp = getAZeroPoint();
2133 if (succeeded(maybeAZp) && verifyAZeroPoint(*maybeAZp).failed())
2134 return failure();
2135
2136 FailureOr<int64_t> maybeBZp = getBZeroPoint();
2137 if (succeeded(maybeBZp) && verifyBZeroPoint(*maybeBZp).failed())
2138 return failure();
2139
2140 return verifyMatMulShapes(*this, /*transposeB=*/false);
2141}
2142
2143LogicalResult tosa::MatMulTOp::inferReturnTypeComponents(
2144 MLIRContext *context, ::std::optional<Location> location,
2145 MatMulTOp::Adaptor adaptor,
2146 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2147 return inferMatMulReturnTypeComponents(ShapeAdaptor(adaptor.getA().getType()),
2148 ShapeAdaptor(adaptor.getB().getType()),
2149 /*transposeB=*/true,
2150 inferredReturnShapes);
2151}
2152
2153LogicalResult MatMulTOp::verify() {
2154 const Type aElementType = getElementTypeOrSelf(getA());
2155 const Type bElementType = getElementTypeOrSelf(getB());
2156
2157 if (failed(
2158 verifyMatMulQuantizedOperandsType(*this, aElementType, bElementType)))
2159 return failure();
2160
2161 if (failed(verifyMatMulZeroPointType(*this, getA(), getAZp(), "a", "a_zp")) ||
2162 failed(verifyMatMulZeroPointType(*this, getB(), getBZp(), "b", "b_zp")))
2163 return failure();
2164
2165 FailureOr<int64_t> maybeAZp = getAZeroPoint();
2166 if (succeeded(maybeAZp) && verifyAZeroPoint(*maybeAZp).failed())
2167 return failure();
2168
2169 FailureOr<int64_t> maybeBZp = getBZeroPoint();
2170 if (succeeded(maybeBZp) && verifyBZeroPoint(*maybeBZp).failed())
2171 return failure();
2172
2173 return verifyMatMulShapes(*this, /*transposeB=*/true);
2174}
2175
2176LogicalResult tosa::MatmulTBlockScaledOp::inferReturnTypeComponents(
2177 MLIRContext *context, ::std::optional<Location> location,
2178 MatmulTBlockScaledOp::Adaptor adaptor,
2179 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2180 SmallVector<int64_t, 3> outShape(3, ShapedType::kDynamic);
2181
2182 const auto aDataShape = cast<ShapedType>(adaptor.getAData().getType());
2183 if (aDataShape.hasRank()) {
2184 outShape[0] = aDataShape.getDimSize(0);
2185 outShape[1] = aDataShape.getDimSize(1);
2186 }
2187
2188 const auto aScaleShape = cast<ShapedType>(adaptor.getAScale().getType());
2189 if (aScaleShape.hasRank()) {
2190 outShape[0] = ShapedType::isDynamic(outShape[0]) ? aScaleShape.getDimSize(0)
2191 : outShape[0];
2192 outShape[1] = ShapedType::isDynamic(outShape[1]) ? aScaleShape.getDimSize(1)
2193 : outShape[1];
2194 }
2195
2196 // If B batch size is 1, it is broadcast across A's batch size
2197 const auto bDataShape = cast<ShapedType>(adaptor.getBData().getType());
2198 if (bDataShape.hasRank()) {
2199 const int64_t bDataBatchSize = bDataShape.getDimSize(0);
2200 if (bDataBatchSize != 1)
2201 outShape[0] =
2202 ShapedType::isDynamic(outShape[0]) ? bDataBatchSize : outShape[0];
2203 outShape[2] = bDataShape.getDimSize(1);
2204 }
2205
2206 const auto bScaleShape = cast<ShapedType>(adaptor.getBScale().getType());
2207 if (bScaleShape.hasRank()) {
2208 const int64_t bScaleBatchSize = bScaleShape.getDimSize(0);
2209 if (bScaleBatchSize != 1)
2210 outShape[0] =
2211 ShapedType::isDynamic(outShape[0]) ? bScaleBatchSize : outShape[0];
2212 outShape[2] = ShapedType::isDynamic(outShape[2]) ? bScaleShape.getDimSize(1)
2213 : outShape[2];
2214 }
2215
2216 inferredReturnShapes.push_back(ShapedTypeComponents(outShape));
2217 return success();
2218}
2219
2220LogicalResult MatmulTBlockScaledOp::verify() {
2221 // Verify same input data types
2222 const Type aDataType = getAData().getType();
2223 const Type bDataType = getBData().getType();
2224 if (failed(verifySameElementTypes(*this, aDataType, bDataType, "A_data",
2225 "B_data")))
2226 return failure();
2227
2228 // Verify input shape compatibility
2229 int64_t N = ShapedType::kDynamic;
2230 int64_t D = ShapedType::kDynamic;
2231 int64_t H = ShapedType::kDynamic;
2232 int64_t W = ShapedType::kDynamic;
2233 int64_t C = ShapedType::kDynamic;
2234 int64_t multiplesOfC = ShapedType::kDynamic;
2235
2236 const ShapeAdaptor aDataShape = ShapeAdaptor(aDataType);
2237 if (aDataShape.hasRank()) {
2238 N = aDataShape.getDimSize(0);
2239 H = aDataShape.getDimSize(1);
2240 C = aDataShape.getDimSize(2);
2241 }
2242
2243 const ShapeAdaptor aScaleShape = ShapeAdaptor(getAScale().getType());
2244 if (aScaleShape.hasRank()) {
2245 if (failed(tryUpdateDimOrFailure(*this, N, aScaleShape.getDimSize(0),
2246 "a_scale", "batch")) ||
2247 failed(tryUpdateDimOrFailure(*this, H, aScaleShape.getDimSize(1),
2248 "a_scale", "height")))
2249 return failure();
2250 multiplesOfC = aScaleShape.getDimSize(2);
2251 }
2252
2253 const ShapeAdaptor bDataShape = ShapeAdaptor(bDataType);
2254 if (bDataShape.hasRank()) {
2255 if (failed(tryUpdateDimOrFailure(*this, D, bDataShape.getDimSize(0),
2256 "b_data", "batch")) ||
2257 failed(tryUpdateDimOrFailure(*this, C, bDataShape.getDimSize(2),
2258 "b_data", "channels")))
2259 return failure();
2260 W = bDataShape.getDimSize(1);
2261 }
2262
2263 const ShapeAdaptor bScaleShape = ShapeAdaptor(getBScale().getType());
2264 if (bScaleShape.hasRank()) {
2265 if (failed(tryUpdateDimOrFailure(*this, D, bScaleShape.getDimSize(0),
2266 "b_scale", "batch")) ||
2267 failed(tryUpdateDimOrFailure(*this, W, bScaleShape.getDimSize(1),
2268 "b_scale", "width")) ||
2269 failed(tryUpdateDimOrFailure(*this, multiplesOfC,
2270 bScaleShape.getDimSize(2), "b_scale",
2271 "C/block_size")))
2272 return failure();
2273 }
2274
2275 // Verify batch size is broadcast compatible
2276 if (ShapedType::isStatic(N) && ShapedType::isStatic(D) && N != D && D != 1)
2277 return emitOpError("expect B matrix batch size to be broadcast compatible "
2278 "with A, got D=")
2279 << D << " vs N=" << N;
2280
2281 // Verify C is a multiple of block size
2282 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(getBlockSize());
2283 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
2284 return emitOpError("expect block size to be 32, got ") << blockSize;
2285 if (ShapedType::isStatic(C) && C % blockSize != 0)
2286 return emitOpError("expect C to be a multiple of block size, got C=")
2287 << C << ", block_size=" << blockSize;
2288
2289 // Verify multiplesOfC is C / block size
2290 if (ShapedType::isStatic(C) && ShapedType::isStatic(multiplesOfC) &&
2291 multiplesOfC != C / blockSize)
2292 return emitOpError(
2293 "expect scale operands dimension 2 to equal C/block_size (")
2294 << C << "/" << blockSize << ")" << ", got " << multiplesOfC;
2295
2296 // Verify output shape
2297 N = ShapedType::isDynamic(N) ? D : N;
2298 const SmallVector<int64_t, 3> expectedOutputShape = {N, H, W};
2299 const auto outputType = cast<ShapedType>(getResult().getType());
2300 if (outputType.hasRank() &&
2301 failed(
2302 verifyCompatibleShape(outputType.getShape(), expectedOutputShape))) {
2303 InFlightDiagnostic opError = emitOpError("expected output shape ");
2304 printShapeToDiagnostic(opError, outputType.getShape());
2305 opError << " to be compatible with expected output shape ";
2306 printShapeToDiagnostic(opError, expectedOutputShape);
2307 return opError;
2308 }
2309
2310 return success();
2311}
2312
2313LogicalResult tosa::PadOp::inferReturnTypeComponents(
2314 MLIRContext *context, ::std::optional<Location> location,
2315 PadOp::Adaptor adaptor,
2316 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2317 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2318 auto paddingRank =
2319 cast<tosa::shapeType>(adaptor.getPadding().getType()).getRank();
2320 SmallVector<int64_t> outputShape;
2321
2322 // If the input rank is unknown, we can infer the output rank using the
2323 // padding shape's rank divided by 2.
2324 if (!inputShape.hasRank()) {
2325 outputShape.resize(paddingRank / 2, ShapedType::kDynamic);
2326 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
2327 return success();
2328 }
2329
2330 SmallVector<int64_t> paddingValues;
2331 // If the paddings value is not a constant, all dimensions must be dynamic.
2332 if (!tosa::getConstShapeValues(adaptor.getPadding().getDefiningOp(),
2333 paddingValues)) {
2334 outputShape.resize(inputShape.getRank(), ShapedType::kDynamic);
2335 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
2336 return success();
2337 }
2338
2339 outputShape.reserve(inputShape.getRank());
2340 for (int i = 0, s = inputShape.getRank(); i < s; i++) {
2341 if (inputShape.isDynamicDim(i)) {
2342 outputShape.push_back(ShapedType::kDynamic);
2343 continue;
2344 }
2345 auto padFront = paddingValues[i * 2];
2346 auto padBack = paddingValues[i * 2 + 1];
2347 if (padFront < 0 || padBack < 0) {
2348 // if either padding for dim i is -1, output dim is unknown
2349 outputShape.push_back(ShapedType::kDynamic);
2350 continue;
2351 }
2352
2353 outputShape.push_back(inputShape.getDimSize(i) + padFront + padBack);
2354 }
2355
2356 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
2357 return success();
2358}
2359
2360LogicalResult tosa::PadOp::verify() {
2361 if (verifySameElementTypes(*this, /* inType = */ getInput1().getType(),
2362 /* outType = */ getOutput().getType())
2363 .failed()) {
2364 return failure();
2365 }
2366
2367 if (auto padConst = getPadConst()) {
2368 if (verifySameElementTypes(*this, /* inType = */ padConst.getType(),
2369 /* outType = */ getOutput().getType())
2370 .failed()) {
2371 return failure();
2372 }
2373 }
2374
2375 RankedTensorType inputType =
2376 llvm::dyn_cast<RankedTensorType>(getInput1().getType());
2377 RankedTensorType outputType =
2378 llvm::dyn_cast<RankedTensorType>(getOutput().getType());
2379 if (!inputType || !outputType)
2380 return success();
2381
2382 if (failed(verifyRanksMatch(getOperation(), inputType, outputType, "input",
2383 "output")))
2384 return failure();
2385
2386 auto inputRank = inputType.getRank();
2387 DenseIntElementsAttr paddingAttr;
2388 if (!matchPattern(getPadding(), m_Constant(&paddingAttr)))
2389 return success();
2390
2391 auto paddingValues = paddingAttr.getValues<APInt>();
2392 if (paddingValues.size() != static_cast<size_t>(inputRank * 2))
2393 return emitOpError() << "padding tensor must have " << inputRank
2394 << " * 2 = " << inputRank * 2 << " elements, but got "
2395 << paddingValues.size();
2396
2397 auto inputShape = inputType.getShape();
2398 auto outputShape = outputType.getShape();
2399
2400 for (int64_t i = 0; i < inputRank; ++i) {
2401 int64_t padStart = paddingValues[i * 2].getSExtValue();
2402 int64_t padEnd = paddingValues[i * 2 + 1].getSExtValue();
2403
2404 if ((padStart < 0 && padStart != -1) || (padEnd < 0 && padEnd != -1)) {
2405 return emitOpError()
2406 << "invalid padding values at dimension " << i
2407 << ": values must be non-negative or -1 for dynamic padding, got ["
2408 << padStart << ", " << padEnd << "]";
2409 }
2410
2411 // Skip shape verification for dynamic input/output
2412 if (inputShape[i] == ShapedType::kDynamic ||
2413 outputShape[i] == ShapedType::kDynamic)
2414 continue;
2415
2416 if (outputShape[i] != inputShape[i] + padStart + padEnd) {
2417 return emitOpError() << "mismatch in output shape at dimension " << i
2418 << ": expected " << inputShape[i] << " + "
2419 << padStart << " + " << padEnd << " = "
2420 << (inputShape[i] + padStart + padEnd)
2421 << ", but got " << outputShape[i];
2422 }
2423 }
2424
2425 return success();
2426}
2427
2428LogicalResult tosa::SliceOp::inferReturnTypeComponents(
2429 MLIRContext *context, ::std::optional<Location> location,
2430 SliceOp::Adaptor adaptor,
2431 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2432
2433 Type inputType = getElementTypeOrSelf(adaptor.getInput1().getType());
2436
2437 if (!tosa::getConstShapeValues(adaptor.getStart().getDefiningOp(), start) ||
2438 !tosa::getConstShapeValues(adaptor.getSize().getDefiningOp(), size)) {
2439 auto rank = cast<tosa::shapeType>(adaptor.getSize().getType()).getRank();
2440 SmallVector<int64_t> fallback(rank, ShapedType::kDynamic);
2441 inferredReturnShapes.push_back(ShapedTypeComponents(fallback, inputType));
2442 return success();
2443 }
2444
2445 // if size[i] is -1, all remaining elements in dimension i are included
2446 // in the slice, similar to TF.
2447 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2448 // initialize outputShape to all unknown
2449 SmallVector<int64_t> outputShape(size.size(), ShapedType::kDynamic);
2450 if (inputShape.hasRank()) {
2451 for (size_t i = 0; i < size.size(); i++) {
2452 if (size[i] != 0 && size[i] >= -1 && start[i] >= 0 &&
2453 (ShapedType::isDynamic(inputShape.getDimSize(i)) ||
2454 start[i] < inputShape.getDimSize(i))) {
2455 // size[i] is not 0 and not < -1, and start[i] is in valid range
2456 if (ShapedType::isDynamic(inputShape.getDimSize(i))) {
2457 // input shape has unknown dim[i] - only valid if size[i] > 0
2458 if (size[i] > 0) {
2459 outputShape[i] = size[i];
2460 }
2461 } else {
2462 // input shape has known dim[i]
2463 if (size[i] == -1) {
2464 outputShape[i] = inputShape.getDimSize(i) - start[i];
2465 } else if (start[i] + size[i] <= inputShape.getDimSize(i)) {
2466 // start[i] + size[i] is within bound of input shape's dim[i]
2467 outputShape[i] = size[i];
2468 }
2469 }
2470 }
2471 }
2472 } else {
2473 outputShape = convertToMlirShape(size);
2474 }
2475 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
2476 return success();
2477}
2478
2479LogicalResult tosa::SliceOp::verify() {
2480 const Value input = getInput1();
2481 const Value output = getOutput();
2482 if (verifySameElementTypes(*this, /* inType = */ input.getType(),
2483 /* outType = */ output.getType())
2484 .failed())
2485 return failure();
2486
2487 const Value start = getStart();
2488 const Value size = getSize();
2489 const ShapeAdaptor inputShape(input.getType());
2490 const ShapeAdaptor outputShape(output.getType());
2491
2492 if (inputShape.hasRank()) {
2493 const auto inputRank = inputShape.getRank();
2494 if (outputShape.hasRank() && inputRank != outputShape.getRank())
2495 return emitOpError(
2496 "expect input1 and output to have the same ranks, got ")
2497 << inputRank << " and " << outputShape.getRank();
2498
2499 const auto startShapeRank =
2500 llvm::cast<tosa::shapeType>(start.getType()).getRank();
2501 if (inputRank != startShapeRank)
2502 return emitOpError("length of start is not equal to rank of input shape");
2503
2504 const auto sizeShapeRank =
2505 llvm::cast<tosa::shapeType>(size.getType()).getRank();
2506 if (inputRank != sizeShapeRank)
2507 return emitOpError("length of size is not equal to rank of input shape");
2508 }
2509
2510 SmallVector<int64_t> startValues;
2511 tosa::getConstShapeValues(start.getDefiningOp(), startValues);
2512 if (startValues.size()) {
2513 if (llvm::any_of(startValues, [](const int64_t v) {
2514 return v < 0 && v != kInferableDimSize;
2515 }))
2516 return emitOpError("start values must be non-negative, got [")
2517 << startValues << "]";
2518 }
2519
2520 // ERROR_IF(is_block_scale<in_out_t>() && start[rank(shape1) - 1] %
2521 // get_innermost_block_size<in_out_t>() != 0);
2522 const auto elemType = getElementTypeOrSelf(input.getType());
2523 if (const auto blockScaledType = llvm::dyn_cast<BlockScaledType>(elemType)) {
2524 const auto startBlock = startValues.back();
2525 const auto scaleBlock =
2526 BlockShapeAttr::getBlockShapeValue(blockScaledType.getBlockShape());
2527 if (startBlock % scaleBlock != 0) {
2528 return emitOpError(
2529 "expected start innermost block size to match data type "
2530 "for block scaled input, got start block=")
2531 << startBlock << ", scale block=" << scaleBlock;
2532 }
2533 }
2534
2535 SmallVector<int64_t> sizeValues;
2536 if (!tosa::getConstShapeValues(size.getDefiningOp(), sizeValues))
2537 return success();
2538
2539 if (llvm::any_of(sizeValues, [](const int64_t v) {
2540 return v <= 0 && v != kInferableDimSize;
2541 }))
2542 return emitOpError("size values must be > 0, got [") << sizeValues << "]";
2543 if (outputShape.hasRank()) {
2544 SmallVector<int64_t> outputDims;
2545 outputShape.getDims(outputDims);
2546 const bool hasNoInferableDims = llvm::all_of(
2547 sizeValues, [](const int64_t v) { return v != kInferableDimSize; });
2548 if (hasNoInferableDims &&
2549 failed(verifyCompatibleShape(outputDims, sizeValues)))
2550 return emitOpError("expected output shape to match size values, got ")
2551 << output.getType() << " vs [" << sizeValues << "]";
2552 }
2553
2554 if (inputShape.hasRank() && startValues.size()) {
2555 SmallVector<int64_t> inputDims;
2556 inputShape.getDims(inputDims);
2557 for (const auto &[index, vals] :
2558 llvm::enumerate(llvm::zip_equal(startValues, sizeValues, inputDims))) {
2559 const auto &[start, size, inputDim] = vals;
2560 if (start == kInferableDimSize || size == kInferableDimSize ||
2561 ShapedType::isDynamic(inputDim))
2562 continue;
2563 if (start + size > inputDim)
2564 return emitOpError("start + size must be less than or equal to input "
2565 "dimension size, got start=")
2566 << start << ", size=" << size
2567 << " vs input dim size=" << inputDim << " at dimension "
2568 << index;
2569 }
2570 }
2571
2572 return success();
2573}
2574
2575LogicalResult tosa::MulOp::inferReturnTypeComponents(
2576 MLIRContext *context, ::std::optional<Location> location,
2577 ValueShapeRange operands, DictionaryAttr attributes, PropertyRef properties,
2578 RegionRange regions,
2579 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2580 // mul op's output shape only depend on input1 and input2, not on shift
2581 ValueShapeRange twoInputs = operands.drop_back();
2583 if (resolveBroadcastShape(twoInputs, outShape).failed()) {
2584 inferredReturnShapes.push_back(ShapedTypeComponents());
2585 } else {
2586 inferredReturnShapes.push_back(ShapedTypeComponents(outShape));
2587 }
2588 return success();
2589}
2590
2591LogicalResult tosa::MulOp::verify() {
2592 const Value output = getOutput();
2593 auto resElemType = getElementTypeOrSelf(output);
2594
2595 // Verify if the element type among operands and result match tosa
2596 // specification.
2597 if (auto resIntType = dyn_cast<IntegerType>(resElemType)) {
2598 IntegerType lhsIntType =
2599 dyn_cast<IntegerType>(getElementTypeOrSelf(getInput1()));
2600 IntegerType rhsIntType =
2601 dyn_cast<IntegerType>(getElementTypeOrSelf(getInput2()));
2602 if (!lhsIntType || !rhsIntType || lhsIntType != rhsIntType)
2603 return emitOpError("requires the same element type for all operands");
2604
2605 // Though the spec requires the element type of result to be i32, a more
2606 // relaxed way is provided at dialect level for easier cooperating with
2607 // other dialects.
2608 if (lhsIntType.getWidth() > resIntType.getWidth())
2609 return emitOpError("invalid data type size for operands or result");
2610
2611 } else {
2612 // For other supported type, the spec requires requires the same element
2613 // type for all operands (excludes `shift` operand) and results.
2614 for (int i = 0; i < 2; ++i) {
2615 if (getElementTypeOrSelf(getOperand(i)) != resElemType)
2616 return emitOpError(
2617 "requires the same element type for all operands and results");
2618 }
2619
2620 // verify shift has value 0 for non-integer types
2621 ElementsAttr shiftElem;
2622 if (matchPattern(getShift(), m_Constant(&shiftElem))) {
2623 int32_t shift = shiftElem.getValues<IntegerAttr>()[0].getInt();
2624 if (shift != 0) {
2625 return emitOpError() << "require shift to be 0 for float type";
2626 }
2627 }
2628 }
2629
2630 // Verify the op has same ranks for all main operands (excludes extra operands
2631 // such as shift of mul op, so this is the only difference with the built-in
2632 // `SameOperandsAndResultRank` trait) and results types, if known.
2633 TypeRange operandTypes = getOperandTypes();
2634 ShapedType aType = cast<ShapedType>(operandTypes[0]);
2635 ShapedType bType = cast<ShapedType>(operandTypes[1]);
2636
2637 const bool aHasRank = aType.hasRank();
2638 const bool bHasRank = bType.hasRank();
2639
2640 bool hasExpectedOutputShape = false;
2641 SmallVector<int64_t> expectedOutputShape;
2642
2643 if (aHasRank && bHasRank) {
2644 const int64_t aRank = aType.getRank();
2645 const int64_t bRank = bType.getRank();
2646 if (aRank != bRank)
2647 return emitOpError("a and b operands don't have matching ranks, got ")
2648 << aRank << " and " << bRank;
2649
2650 // check for broadcast compatible shapes
2652 aType.getShape(), bType.getShape(), expectedOutputShape))
2653 return emitOpError("a and b operands don't have broadcast-compatible "
2654 "shapes, got ")
2655 << aType << " and " << bType;
2656 hasExpectedOutputShape = true;
2657 }
2658
2659 ShapedType resultType = cast<ShapedType>(output.getType());
2660 if (!resultType.hasRank())
2661 return success();
2662
2663 const int64_t resultRank = resultType.getRank();
2664 if (aHasRank && resultRank != aType.getRank())
2665 return emitOpError("result type has different rank than a, got ")
2666 << resultRank << " vs " << aType.getRank();
2667 if (bHasRank && resultRank != bType.getRank())
2668 return emitOpError("result type has different rank than b, got ")
2669 << resultRank << " vs " << bType.getRank();
2670
2671 if (hasExpectedOutputShape &&
2672 failed(verifyOutputShapeCompatibleWithExpected(getOperation(), resultType,
2673 expectedOutputShape)))
2674 return failure();
2675
2676 return success();
2677}
2678
2679LogicalResult tosa::TableOp::inferReturnTypeComponents(
2680 MLIRContext *context, ::std::optional<Location> location,
2681 TableOp::Adaptor adaptor,
2682 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2683 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2684
2685 if (!inputShape.hasRank()) {
2686 inferredReturnShapes.push_back(ShapedTypeComponents());
2687 return success();
2688 }
2689
2690 inferredReturnShapes.resize(1);
2691 inputShape.getDims(inferredReturnShapes[0]);
2692 return success();
2693}
2694
2695LogicalResult tosa::TableOp::verify() {
2696 const TensorType inputType = getInput1().getType();
2697 const TensorType outputType = getOutput().getType();
2698
2699 if (!inputType.hasRank() || !outputType.hasRank())
2700 return success();
2701
2702 if (failed(verifyRanksMatch(getOperation(), inputType, outputType, "input",
2703 "result")))
2704 return failure();
2705
2706 auto inputDims = inputType.getShape();
2707 auto outputDims = outputType.getShape();
2708 for (auto it : llvm::enumerate(llvm::zip(inputDims, outputDims))) {
2709 int64_t dim = it.index();
2710 auto [inputDim, outputDim] = it.value();
2711 if (ShapedType::isStatic(outputDim) && outputDim != inputDim) {
2712 return emitOpError() << "dim(result, " << dim << ") = " << outputDim
2713 << " doesn't match dim(input, " << dim
2714 << ") = " << inputDim;
2715 }
2716 }
2717 return success();
2718}
2719
2720LogicalResult
2721tosa::TileOp::getConstantMultiples(SmallVector<int64_t> &multiples) {
2722 // Multiples must be constants.
2723 DenseIntElementsAttr multiplesAttr;
2724 if (!matchPattern(getMultiples(), m_Constant(&multiplesAttr)))
2725 return failure();
2726 multiples =
2727 llvm::map_to_vector(multiplesAttr.getValues<APInt>(),
2728 [](const APInt &val) { return val.getSExtValue(); });
2729 return success();
2730}
2731
2732LogicalResult tosa::TileOp::inferReturnTypeComponents(
2733 MLIRContext *context, ::std::optional<Location> location,
2734 TileOp::Adaptor adaptor,
2735 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2736 Type inputType = getElementTypeOrSelf(adaptor.getInput1().getType());
2737 SmallVector<int64_t> multiples;
2738 if (!tosa::getConstShapeValues(adaptor.getMultiples().getDefiningOp(),
2739 multiples)) {
2740 auto rank =
2741 cast<tosa::shapeType>(adaptor.getMultiples().getType()).getRank();
2742 SmallVector<int64_t> fallback(rank, ShapedType::kDynamic);
2743 inferredReturnShapes.push_back(ShapedTypeComponents(fallback, inputType));
2744 return success();
2745 }
2746 multiples = convertToMlirShape(multiples);
2747
2748 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2749 SmallVector<int64_t> outputShape;
2750 if (!inputShape.hasRank()) {
2751 outputShape.resize(multiples.size(), ShapedType::kDynamic);
2752 inferredReturnShapes.push_back(
2753 ShapedTypeComponents(outputShape, inputType));
2754 return success();
2755 }
2756 if (static_cast<size_t>(inputShape.getRank()) != multiples.size())
2757 return failure();
2758
2759 // Any non dynamic dimension can be multiplied to a known size.
2760 outputShape.reserve(multiples.size());
2761 for (int i = 0, s = inputShape.getRank(); i < s; i++) {
2762 if (multiples[i] == ShapedType::kDynamic) {
2763 outputShape.push_back(ShapedType::kDynamic);
2764 } else {
2765 int64_t dim = inputShape.getDimSize(i);
2766 if (dim != ShapedType::kDynamic)
2767 dim *= multiples[i];
2768 outputShape.push_back(dim);
2769 }
2770 }
2771
2772 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape, inputType));
2773 return success();
2774}
2775
2776LogicalResult tosa::TileOp::verify() {
2777 if (verifySameElementTypes(*this, /* intype = */ getInput1().getType(),
2778 /* outType = */ getOutput().getType())
2779 .failed()) {
2780 return failure();
2781 }
2782 ShapedType inputType = llvm::cast<ShapedType>(getInput1().getType());
2783 ShapedType outputType = llvm::cast<ShapedType>(getType());
2784
2785 shapeType multiplesType =
2786 llvm::cast<tosa::shapeType>(getMultiples().getType());
2787
2788 auto multiplesRank = multiplesType.getRank();
2789
2790 if (inputType.hasRank()) {
2791 if (inputType.getRank() != multiplesRank)
2792 return emitOpError("expect 'multiples' to have rank ")
2793 << inputType.getRank() << " but got " << multiplesRank << ".";
2794 if (outputType.hasRank() &&
2795 failed(verifyRanksMatch(getOperation(), inputType, outputType, "input",
2796 "output")))
2797 return failure();
2798 } else if (outputType.hasRank() && outputType.getRank() != multiplesRank)
2799 return emitOpError("expect 'multiples' array to have length ")
2800 << outputType.getRank() << " but got " << multiplesRank << ".";
2801
2802 SmallVector<int64_t> multiples;
2803 if (getConstantMultiples(multiples).succeeded() &&
2804 llvm::any_of(multiples, [](int64_t v) { return v <= 0 && v != -1; }))
2805 return emitOpError(
2806 "expect element of 'multiples' to be positive integer or -1.");
2807
2808 return success();
2809}
2810
2811bool tosa::ReshapeOp::isCompatibleReturnTypes(TypeRange l, TypeRange r) {
2812 if (l.size() != r.size() || l.size() != 1)
2813 return false;
2814 return getElementTypeOrSelf(l[0]) == getElementTypeOrSelf(r[0]);
2815}
2816
2817LogicalResult tosa::ReshapeOp::inferReturnTypeComponents(
2818 MLIRContext *context, ::std::optional<Location> location,
2819 ReshapeOp::Adaptor adaptor,
2820 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2821 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2822 Type inputType = getElementTypeOrSelf(adaptor.getInput1().getType());
2823 llvm::SmallVector<int64_t> newShapeValue;
2824 if (!tosa::getConstShapeValues(adaptor.getShape().getDefiningOp(),
2825 newShapeValue)) {
2826 auto rank = cast<tosa::shapeType>(adaptor.getShape().getType()).getRank();
2827 SmallVector<int64_t> fallback(rank, ShapedType::kDynamic);
2828 inferredReturnShapes.push_back(ShapedTypeComponents(fallback, inputType));
2829 return success();
2830 }
2831 newShapeValue = convertToMlirShape(newShapeValue);
2832
2833 // We cannot infer from the total number of elements so we must take the
2834 // shape attribute as exact.
2835 if (!inputShape.hasRank() || !inputShape.hasStaticShape()) {
2836 inferredReturnShapes.push_back(
2837 ShapedTypeComponents(newShapeValue, inputType));
2838 return success();
2839 }
2840
2841 // Determine the number of elements covered by the slice of all static
2842 // dimensions. This allows us to infer the length of the remaining dynamic
2843 // dimension.
2844 int64_t numElements = inputShape.getNumElements();
2845 int64_t staticMul = 1;
2846 for (auto val : newShapeValue) {
2847 if (ShapedType::isStatic(val)) {
2848 staticMul *= val;
2849 }
2850 }
2851
2852 // Determine the length of the dynamic dimension.
2853 for (auto &val : newShapeValue) {
2854 if (ShapedType::isDynamic(val))
2855 val = numElements / staticMul;
2856 }
2857
2858 inferredReturnShapes.push_back(
2859 ShapedTypeComponents(newShapeValue, inputType));
2860 return success();
2861}
2862
2863llvm::LogicalResult tosa::ReshapeOp::verify() {
2864 if (verifySameElementTypes(*this, /* inType = */ getInput1().getType(),
2865 /* outType = */ getOutput().getType())
2866 .failed()) {
2867 return failure();
2868 }
2869 TensorType inputType = getInput1().getType();
2870
2871 SmallVector<int64_t> shapeValues;
2872 if (!tosa::getConstShapeValues(getShape().getDefiningOp(), shapeValues)) {
2873 // skip following checks if shape is not constant
2874 return mlir::success();
2875 }
2876
2877 int missingDims = llvm::count(shapeValues, kInferableDimSize);
2878 if (missingDims > 1)
2879 return emitOpError() << "expected at most one target dimension to be "
2881
2882 const auto outputType = dyn_cast<RankedTensorType>(getType());
2883 if (!outputType)
2884 return success();
2885
2886 if ((int64_t)shapeValues.size() != outputType.getRank())
2887 return emitOpError() << "new shape does not match result rank";
2888
2889 for (auto [newShapeDim, outputShapeDim] :
2890 zip(shapeValues, outputType.getShape())) {
2891 if (newShapeDim != kInferableDimSize &&
2892 newShapeDim != ShapedType::kDynamic &&
2893 outputShapeDim != ShapedType::kDynamic && newShapeDim != outputShapeDim)
2894 return emitOpError() << "new shape is inconsistent with result shape";
2895
2896 if (newShapeDim != ShapedType::kDynamic && newShapeDim < kInferableDimSize)
2897 return emitOpError() << "new shape has invalid tensor dimension size "
2898 << newShapeDim;
2899 }
2900
2901 if (inputType.hasStaticShape()) {
2902 int64_t inputElementsNum = inputType.getNumElements();
2903 if (outputType.hasStaticShape()) {
2904 int64_t outputElementsNum = outputType.getNumElements();
2905 if (inputElementsNum != outputElementsNum) {
2906 return emitOpError() << "cannot reshape " << inputElementsNum
2907 << " elements into " << outputElementsNum;
2908 }
2909 }
2910
2911 int64_t newShapeElementsNum =
2912 llvm::accumulate(shapeValues, int64_t(1), [](int64_t acc, int64_t dim) {
2913 return (dim > 0) ? acc * dim : acc;
2914 });
2915 bool isStaticNewShape =
2916 llvm::all_of(shapeValues, [](int64_t s) { return s > 0; });
2917 if ((isStaticNewShape && inputElementsNum != newShapeElementsNum) ||
2918 (!isStaticNewShape && newShapeElementsNum > inputElementsNum)) {
2919 return emitOpError() << "cannot reshape " << inputElementsNum
2920 << " elements into " << newShapeElementsNum;
2921 }
2922 }
2923
2924 return mlir::success();
2925}
2926
2927bool tosa::ReshapeBlockScaledOp::isCompatibleReturnTypes(TypeRange l,
2928 TypeRange r) {
2929 if (l.size() != r.size() || l.size() < 1 || l.size() > 2)
2930 return false;
2931 bool ok = (getElementTypeOrSelf(l[0]) == getElementTypeOrSelf(r[0]));
2932 if (l.size() == 2)
2933 ok = ok && (getElementTypeOrSelf(l[1]) == getElementTypeOrSelf(r[1]));
2934 return ok;
2935}
2936
2937LogicalResult tosa::ReshapeBlockScaledOp::inferReturnTypeComponents(
2938 MLIRContext *context, ::std::optional<Location> location,
2939 ReshapeBlockScaledOp::Adaptor adaptor,
2940 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
2941
2942 const auto numInputs = adaptor.getInput().size();
2943 ShapeAdaptor inputShape(adaptor.getInput()[0].getType());
2944 Type inputType = getElementTypeOrSelf(adaptor.getInput()[0].getType());
2945 llvm::SmallVector<int64_t> newShapeValue;
2946 const auto newShape = adaptor.getNewValueShape();
2947 if (!tosa::getConstShapeValues(newShape.getDefiningOp(), newShapeValue)) {
2948 auto rank = cast<tosa::shapeType>(newShape.getType()).getRank();
2949 SmallVector<int64_t> fallback(rank, ShapedType::kDynamic);
2950 inferredReturnShapes.push_back(ShapedTypeComponents(fallback, inputType));
2951 if (numInputs == 2)
2952 inferredReturnShapes.push_back(ShapedTypeComponents(
2953 fallback, getElementTypeOrSelf(adaptor.getInput()[1].getType())));
2954 return success();
2955 }
2956
2957 const uint32_t blockSize =
2958 BlockSizeAttr::getBlockSizeValue(adaptor.getBlockSize());
2959
2960 llvm::SmallVector<int64_t> newScaleShapeValue;
2961 if (numInputs == 2) {
2962 newScaleShapeValue.assign(newShapeValue.begin(), newShapeValue.end());
2963 if (!newScaleShapeValue.empty() &&
2964 ShapedType::isStatic(newScaleShapeValue.back()))
2965 newScaleShapeValue.back() /= blockSize;
2966 }
2967
2968 inferredReturnShapes.push_back(
2969 ShapedTypeComponents(newShapeValue, inputType));
2970 if (numInputs == 2) {
2971 // Fix up scale shape - with special case for last dimension
2972 for (size_t idx = 0; idx < newShapeValue.size(); idx++) {
2973 if (ShapedType::isDynamic(newScaleShapeValue[idx])) {
2974 newScaleShapeValue[idx] = newShapeValue[idx];
2975 if (idx + 1 == newShapeValue.size())
2976 newScaleShapeValue[idx] /= blockSize;
2977 }
2978 }
2979
2980 inferredReturnShapes.push_back(ShapedTypeComponents(
2981 newScaleShapeValue,
2982 getElementTypeOrSelf(adaptor.getInput()[1].getType())));
2983 }
2984 return success();
2985}
2986
2987llvm::LogicalResult tosa::ReshapeBlockScaledOp::verify() {
2988 const Operation::operand_range inputList = getInput();
2989 const Operation::result_range outputList = getResults();
2990
2991 if (inputList.size() == 0)
2992 return emitOpError("requires at least one input");
2993
2994 if (inputList.size() > 2)
2995 return emitOpError("requires at most two inputs");
2996
2997 if (inputList.size() != outputList.size())
2998 return emitOpError("requires number of results to match inputs");
2999
3000 if (verifySameElementTypes(*this, /* inType = */ inputList[0].getType(),
3001 /* outType = */ outputList[0].getType())
3002 .failed()) {
3003 return failure();
3004 }
3005
3006 if (inputList.size() == 2 &&
3007 cast<tosa::shapeType>(getNewValueShape().getType()).getRank() == 0)
3008 return emitOpError("requires new shape to have a rank greater than 0");
3009
3010 const auto inputType = llvm::cast<ShapedType>(inputList[0].getType());
3011 if (!inputType.hasRank())
3012 return success();
3013 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(getBlockSize());
3014
3015 if (inputList.size() == 2) {
3016 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
3017 return emitOpError("expect block size to be 32, got ") << blockSize;
3018 if (llvm::any_of(inputList, [](Value v) {
3019 const auto input = cast<ShapedType>(v.getType());
3020 return input.hasRank() && input.getRank() == 0;
3021 }))
3022 return emitOpError(
3023 "requires all input shapes have a rank greater than 0");
3024 if (llvm::any_of(outputList, [](Value v) {
3025 const auto output = cast<ShapedType>(v.getType());
3026 return output.hasRank() && output.getRank() == 0;
3027 }))
3028 return emitOpError(
3029 "requires all result shapes have a rank greater than 0");
3030
3031 if (verifySameElementTypes(*this, /* inType = */ inputList[1].getType(),
3032 /* outType = */ outputList[1].getType())
3033 .failed()) {
3034 return failure();
3035 }
3036
3037 const auto inputScaleType = llvm::cast<ShapedType>(inputList[1].getType());
3038 if (inputScaleType.hasRank()) {
3039 if (inputType.getRank() != inputScaleType.getRank())
3040 return emitOpError("input shapes do not have same rank");
3041
3042 // Check all but the last dimension that the input shape dimensions match
3043 for (auto dimIdx = 0; dimIdx < inputType.getRank() - 1; dimIdx++) {
3044 const int64_t inputValueDim = inputType.getDimSize(dimIdx);
3045 const int64_t inputScaleDim = inputScaleType.getShape()[dimIdx];
3046 if (ShapedType::isStatic(inputValueDim) &&
3047 ShapedType::isStatic(inputScaleDim) &&
3048 inputValueDim != inputScaleDim)
3049 return emitOpError("input shapes for data and scale do not match on "
3050 "dimension ")
3051 << dimIdx;
3052 }
3053
3054 // Verify last dimension of input is a multiple of block size
3055 const int64_t lastValueDim =
3056 inputType.getDimSize(inputType.getRank() - 1);
3057 if (ShapedType::isStatic(lastValueDim)) {
3058 if (lastValueDim % blockSize != 0)
3059 return emitOpError("expect last dimension of input_data (")
3060 << lastValueDim << ") to be divisible by block_size ("
3061 << blockSize << ")";
3062
3063 const int64_t lastScaleDim =
3064 inputScaleType.getDimSize(inputScaleType.getRank() - 1);
3065 // Verify last dimension of scale is lastValueDim / block size
3066 if (ShapedType::isStatic(lastScaleDim) &&
3067 lastScaleDim != lastValueDim / blockSize)
3068 return emitOpError("expect last dimension of scale_data (")
3069 << lastScaleDim << ") to be " << lastValueDim << "/"
3070 << blockSize;
3071 }
3072 }
3073 } else {
3074 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_1))
3075 return emitOpError("expect block size to be 1, got ") << blockSize;
3076 }
3077
3078 // Get the new value shape dimension values.
3079 SmallVector<int64_t> shapeValues;
3080 if (!tosa::getConstShapeValues(getNewValueShape().getDefiningOp(),
3081 shapeValues)) {
3082 // skip following checks if shape is not constant
3083 return mlir::success();
3084 }
3085
3086 if (inputList.size() == 2) {
3087 const int64_t lastShapeDim = shapeValues.back();
3088 if (ShapedType::isStatic(lastShapeDim) && lastShapeDim % blockSize != 0)
3089 return emitOpError("expect last dimension of new shape (")
3090 << lastShapeDim << ") to be divisible by block_size (" << blockSize
3091 << ")";
3092 }
3093
3094 const auto outputType = llvm::cast<ShapedType>(outputList[0].getType());
3095 if (!outputType.hasRank())
3096 return success();
3097
3098 if (static_cast<int64_t>(shapeValues.size()) != outputType.getRank())
3099 return emitOpError() << "result does not match new shape rank";
3100
3101 for (auto [newShapeDim, outputShapeDim] :
3102 zip(shapeValues, outputType.getShape())) {
3103 if (ShapedType::isStatic(newShapeDim) &&
3104 ShapedType::isStatic(outputShapeDim) && newShapeDim != outputShapeDim)
3105 return emitOpError() << "result shape is inconsistent with new shape";
3106 }
3107
3108 if (outputList.size() == 2) {
3109 // Set up scale shape from new shape given
3110 SmallVector<int64_t> scaleShapeValues(shapeValues.begin(),
3111 shapeValues.end());
3112 scaleShapeValues.back() /= blockSize;
3113
3114 const auto outputScaleType =
3115 llvm::cast<ShapedType>(outputList[1].getType());
3116 if (outputScaleType.hasRank()) {
3117 if ((int64_t)scaleShapeValues.size() != outputScaleType.getRank())
3118 return emitOpError() << "result scale does not match new shape rank";
3119
3120 for (auto [newScaleShapeDim, outputScaleShapeDim] :
3121 zip(scaleShapeValues, outputScaleType.getShape())) {
3122 if (ShapedType::isStatic(newScaleShapeDim) &&
3123 ShapedType::isStatic(outputScaleShapeDim) &&
3124 newScaleShapeDim != outputScaleShapeDim)
3125 return emitOpError()
3126 << "result scale shape is inconsistent with new shape";
3127 }
3128 }
3129 }
3130
3131 if (inputType.hasStaticShape()) {
3132 int64_t inputElementsNum = inputType.getNumElements();
3133 if (outputType.hasStaticShape()) {
3134 int64_t outputElementsNum = outputType.getNumElements();
3135 if (inputElementsNum != outputElementsNum) {
3136 return emitOpError() << "cannot reshape " << inputElementsNum
3137 << " elements into " << outputElementsNum;
3138 }
3139 }
3140
3141 int64_t newShapeElementsNum =
3142 llvm::accumulate(shapeValues, int64_t(1), [](int64_t acc, int64_t dim) {
3143 return (dim > 0) ? acc * dim : acc;
3144 });
3145 bool isStaticNewShape =
3146 llvm::all_of(shapeValues, [](int64_t s) { return s > 0; });
3147 if ((isStaticNewShape && inputElementsNum != newShapeElementsNum) ||
3148 (!isStaticNewShape && newShapeElementsNum > inputElementsNum)) {
3149 return emitOpError() << "cannot reshape " << inputElementsNum
3150 << " elements into " << newShapeElementsNum;
3151 }
3152 }
3153
3154 return mlir::success();
3155}
3156
3157// return failure if val is not a constant
3158// set zp to -1 if val is non-zero float or val is not integer nor float
3159// otherwise set zp to val's constant value
3160static FailureOr<int64_t> getZeroPoint(Value val, bool signExtend) {
3161 ElementsAttr zpAttr;
3162 if (!matchPattern(val, m_Constant(&zpAttr))) {
3163 return failure();
3164 }
3165
3166 Type zpElemType = zpAttr.getElementType();
3167
3168 if (llvm::isa<FloatType>(zpElemType)) {
3169 if (zpAttr.getValues<APFloat>()[0].isZero()) {
3170 return 0;
3171 }
3172 // return non-zero value to trigger error check
3173 return -1;
3174 }
3175
3176 if (llvm::isa<IntegerType>(zpElemType)) {
3177 if (signExtend)
3178 return zpAttr.getValues<APInt>()[0].getSExtValue();
3179 return zpAttr.getValues<APInt>()[0].getZExtValue();
3180 }
3181
3182 // return non-zero value to trigger error check
3183 return -1;
3184}
3185
3186template <typename T>
3187static LogicalResult verifyZeroPoint(T op, Value val, const int64_t &zp,
3188 const std::string &operand) {
3189 Type zpElemType = getElementTypeOrSelf(val);
3190
3191 if (!zpElemType.isInteger(8) && zp != 0) {
3192 // convert operand to lower case for error message
3193 std::string lower = operand;
3194 llvm::transform(lower, lower.begin(), ::tolower);
3195 return op.emitOpError()
3196 << lower << " zero point must be zero for non-int8 integer types";
3197 }
3198
3199 return success();
3200}
3201
3202static LogicalResult verifyZeroPoint(tosa::RescaleOp op, Value zpVal,
3203 const int64_t &zp,
3204 const std::string &operand) {
3205 bool isInputZp = (operand == "Input");
3206
3207 bool tensorUnsigned =
3208 isInputZp ? op.getInputUnsigned() : op.getOutputUnsigned();
3209 StringRef tensorName = isInputZp ? "input" : "output";
3210
3211 Type zpElemType = getElementTypeOrSelf(zpVal);
3212
3213 if (zp != 0) {
3214 if (!zpElemType.isInteger(8) &&
3215 !(zpElemType.isInteger(16) && tensorUnsigned)) {
3216 return op.emitOpError()
3217 << "expect " << tensorName << "_zp of 0, got " << zp;
3218 }
3219 if (zpElemType.isInteger(16) && tensorUnsigned && zp != 32768) {
3220 return op.emitOpError() << "expect " << tensorName
3221 << "_zp of 0 or 32768 for unsigned int16 "
3222 << tensorName << ", got " << zp;
3223 }
3224 }
3225
3226 return success();
3227}
3228
3229#define ZERO_POINT_HELPER(OP, OPERAND_NAME, SIGN_EXTEND) \
3230 FailureOr<int64_t> tosa::OP::get##OPERAND_NAME##ZeroPoint() { \
3231 return getZeroPoint(get##OPERAND_NAME##Zp(), SIGN_EXTEND); \
3232 } \
3233 LogicalResult tosa::OP::verify##OPERAND_NAME##ZeroPoint(int64_t zp) { \
3234 return verifyZeroPoint(*this, get##OPERAND_NAME##Zp(), zp, #OPERAND_NAME); \
3235 }
3236
3237ZERO_POINT_HELPER(Conv2DOp, Input, true)
3238ZERO_POINT_HELPER(Conv2DOp, Weight, true)
3239ZERO_POINT_HELPER(Conv3DOp, Input, true)
3240ZERO_POINT_HELPER(Conv3DOp, Weight, true)
3241ZERO_POINT_HELPER(DepthwiseConv2DOp, Input, true)
3242ZERO_POINT_HELPER(DepthwiseConv2DOp, Weight, true)
3243ZERO_POINT_HELPER(TransposeConv2DOp, Input, true)
3244ZERO_POINT_HELPER(TransposeConv2DOp, Weight, true)
3245ZERO_POINT_HELPER(AvgPool2dOp, Input, true)
3246ZERO_POINT_HELPER(AvgPool2dOp, Output, true)
3247ZERO_POINT_HELPER(AvgPool2dAdaptiveOp, Input, true)
3248ZERO_POINT_HELPER(AvgPool2dAdaptiveOp, Output, true)
3249ZERO_POINT_HELPER(MatMulOp, A, true)
3250ZERO_POINT_HELPER(MatMulOp, B, true)
3251ZERO_POINT_HELPER(MatMulTOp, A, true)
3252ZERO_POINT_HELPER(MatMulTOp, B, true)
3253ZERO_POINT_HELPER(NegateOp, Input1, true)
3254ZERO_POINT_HELPER(NegateOp, Output, true)
3255ZERO_POINT_HELPER(RescaleOp, Input, !getInputUnsigned())
3256ZERO_POINT_HELPER(RescaleOp, Output, !getOutputUnsigned())
3257#undef ZERO_POINT_HELPER
3258
3259LogicalResult tosa::TransposeOp::inferReturnTypeComponents(
3260 MLIRContext *context, ::std::optional<Location> location,
3261 TransposeOp::Adaptor adaptor,
3262 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3263 ShapeAdaptor inputShape(adaptor.getInput1().getType());
3264
3265 // If input rank and permutation length is unknown, the output rank is
3266 // unknown.
3267 if (!inputShape.hasRank()) {
3268 inferredReturnShapes.push_back(ShapedTypeComponents());
3269 return success();
3270 }
3271
3272 const auto inputRank = inputShape.getRank();
3273
3274 // This would imply the number of permutations does not match the rank of
3275 // the input which is illegal.
3276 if (adaptor.getPerms().size() != static_cast<size_t>(inputRank)) {
3277 return failure();
3278 }
3279
3280 SmallVector<int64_t> outputShape;
3281 // Rank-0 means no permutations matter.
3282 if (inputRank == 0) {
3283 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3284 return success();
3285 }
3286
3287 // Check whether the input dimensions are all the same.
3288 bool allTheSame = true;
3289 for (int i = 1, s = inputRank; i < s; i++) {
3290 if (inputShape.getDimSize(0) != inputShape.getDimSize(i)) {
3291 allTheSame = false;
3292 break;
3293 }
3294 }
3295
3296 // If all of the input dimensions are the same we don't care about the
3297 // permutation.
3298 if (allTheSame) {
3299 outputShape.resize(inputRank, inputShape.getDimSize(0));
3300 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3301 return success();
3302 }
3303
3304 outputShape.resize(inputRank, ShapedType::kDynamic);
3305
3306 // Constant permutation values must be within the input rank.
3307 if (llvm::any_of(adaptor.getPerms(),
3308 [inputRank](const auto i) { return i >= inputRank; }))
3309 return failure();
3310
3311 outputShape.reserve(inputRank);
3312 for (int i = 0, s = inputRank; i < s; i++) {
3313 outputShape[i] = inputShape.getDimSize(adaptor.getPerms()[i]);
3314 }
3315
3316 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3317 return success();
3318}
3319
3320LogicalResult tosa::TransposeOp::verify() {
3321 if (verifySameElementTypes(*this, /* inType = */ getInput1().getType(),
3322 /* outType = */ getOutput().getType())
3323 .failed()) {
3324 return failure();
3325 }
3326
3327 const ShapeAdaptor inputShape(getInput1().getType());
3328 const ShapeAdaptor outputShape(getOutput().getType());
3329
3330 const llvm::ArrayRef<int32_t> constantPerms = getPerms();
3331
3332 if (inputShape.hasRank() &&
3333 constantPerms.size() != static_cast<size_t>(inputShape.getRank()))
3334 return emitOpError() << "expected perms attribute to have size "
3335 << inputShape.getRank()
3336 << " (input rank) but got size "
3337 << constantPerms.size();
3338
3339 if (inputShape.hasRank() && outputShape.hasRank() &&
3340 inputShape.getRank() != outputShape.getRank())
3341 return emitOpError()
3342 << "expected input tensor rank to equal result tensor rank";
3343
3344 if (outputShape.hasRank() &&
3345 constantPerms.size() != static_cast<size_t>(outputShape.getRank()))
3346 return emitOpError() << "expected perms attribute to have size "
3347 << outputShape.getRank()
3348 << " (output rank) but got size "
3349 << constantPerms.size();
3350
3351 if (!llvm::all_of(constantPerms,
3352 [&constantPerms](int32_t s) {
3353 return s >= 0 &&
3354 static_cast<size_t>(s) < constantPerms.size();
3355 }) ||
3356 !isPermutationVector(llvm::map_to_vector(
3357 constantPerms, [](int32_t v) -> int64_t { return v; })))
3358 return emitOpError() << "expected valid permutation indices";
3359
3360 if (isa<BlockScaledType>(getInput1().getType().getElementType()) &&
3361 constantPerms.back() != static_cast<int32_t>(constantPerms.size()) - 1) {
3362 return emitOpError() << "expected no-op permutation on innermost dimension "
3363 "for block scaled input";
3364 }
3365
3366 // ERROR_IF(tensor_size(shape1) != tensor_size(shape))
3367 if (inputShape.hasStaticShape() && outputShape.hasStaticShape() &&
3368 inputShape.getNumElements() != outputShape.getNumElements())
3369 return emitOpError() << "expected input1 and output to have same numbers "
3370 "of elements, got "
3371 << inputShape.getNumElements() << " and "
3372 << outputShape.getNumElements();
3373
3374 // Verify that the types of the input and output tensors are properly
3375 // permuted.
3376 if (inputShape.hasRank() && outputShape.hasRank()) {
3377 for (auto i = 0; i < outputShape.getRank(); i++) {
3378 if (inputShape.isDynamicDim(constantPerms[i]) ||
3379 outputShape.isDynamicDim(i))
3380 continue;
3381
3382 if (inputShape.getDimSize(constantPerms[i]) != outputShape.getDimSize(i))
3383 return emitOpError()
3384 << "expected output tensor dim " << i << " to match "
3385 << "input dim " << constantPerms[i] << " with value of "
3386 << inputShape.getDimSize(constantPerms[i]);
3387 }
3388 }
3389
3390 return success();
3391}
3392
3393LogicalResult TransposeOp::reifyResultShapes(
3394 OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes) {
3395
3396 const llvm::ArrayRef<int32_t> transposePerms = getPerms();
3397
3398 Value input = getInput1();
3399 auto inputType = cast<TensorType>(input.getType());
3400
3401 SmallVector<OpFoldResult> returnedDims(inputType.getRank());
3402 for (auto dim : transposePerms) {
3403 int32_t dimInInput = transposePerms[dim];
3404 if (inputType.isDynamicDim(dimInInput))
3405 returnedDims[dim] =
3406 tensor::DimOp::create(builder, getLoc(), input, dimInInput)
3407 .getResult();
3408 else
3409 returnedDims[dim] =
3410 builder.getIndexAttr(inputType.getDimSize(dimInInput));
3411 }
3412
3413 reifiedReturnShapes.emplace_back(std::move(returnedDims));
3414 return success();
3415}
3416
3417LogicalResult tosa::GatherOp::inferReturnTypeComponents(
3418 MLIRContext *context, ::std::optional<Location> location,
3419 GatherOp::Adaptor adaptor,
3420 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3421 llvm::SmallVector<int64_t> outputShape;
3422 outputShape.resize(3, ShapedType::kDynamic);
3423
3424 ShapeAdaptor valuesShape(adaptor.getValues().getType());
3425 if (valuesShape.hasRank()) {
3426 outputShape[0] = valuesShape.getDimSize(0);
3427 outputShape[2] = valuesShape.getDimSize(2);
3428 }
3429
3430 ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3431 if (indicesShape.hasRank()) {
3432 if (outputShape[0] == ShapedType::kDynamic)
3433 outputShape[0] = indicesShape.getDimSize(0);
3434 if (outputShape[1] == ShapedType::kDynamic)
3435 outputShape[1] = indicesShape.getDimSize(1);
3436 }
3437
3438 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3439 return success();
3440}
3441
3442LogicalResult tosa::RowGatherOp::inferReturnTypeComponents(
3443 MLIRContext *context, ::std::optional<Location> location,
3444 RowGatherOp::Adaptor adaptor,
3445 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3446 llvm::SmallVector<int64_t> outputShape;
3447 outputShape.resize(3, ShapedType::kDynamic);
3448
3449 const ShapeAdaptor valuesShape(adaptor.getValues().getType());
3450 if (valuesShape.hasRank()) {
3451 outputShape[0] = valuesShape.getDimSize(0);
3452 outputShape[2] = valuesShape.getDimSize(2);
3453 }
3454
3455 const ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3456 if (indicesShape.hasRank()) {
3457 if (outputShape[0] == ShapedType::kDynamic)
3458 outputShape[0] = indicesShape.getDimSize(0);
3459
3460 const FailureOr<int32_t> maybeRowCount =
3461 getConstantScalarIntValue<int32_t>(adaptor.getRowCount());
3462 if (succeeded(maybeRowCount)) {
3463 const int64_t indicesW = indicesShape.getDimSize(1);
3464 if (ShapedType::isStatic(indicesW))
3465 outputShape[1] = indicesW * maybeRowCount.value();
3466 }
3467 }
3468
3469 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3470 return success();
3471}
3472
3473LogicalResult tosa::RowGatherBlockScaledOp::inferReturnTypeComponents(
3474 MLIRContext *context, ::std::optional<Location> location,
3475 RowGatherBlockScaledOp::Adaptor adaptor,
3476 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3477 const auto values = adaptor.getValues();
3478 if (values.empty())
3479 return failure();
3480
3481 SmallVector<int64_t> dataShape(3, ShapedType::kDynamic);
3482 const ShapeAdaptor valuesShape(values.front().getType());
3483 if (valuesShape.hasRank()) {
3484 dataShape[0] = valuesShape.getDimSize(0);
3485 dataShape[2] = valuesShape.getDimSize(2);
3486 }
3487
3488 const ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3489 if (indicesShape.hasRank()) {
3490 if (dataShape[0] == ShapedType::kDynamic)
3491 dataShape[0] = indicesShape.getDimSize(0);
3492
3493 if (auto rowCount =
3494 getConstantScalarIntValue<int32_t>(adaptor.getRowCount());
3495 succeeded(rowCount) && rowCount.value() > 0) {
3496 const int64_t indicesW = indicesShape.getDimSize(1);
3497 if (ShapedType::isStatic(indicesW))
3498 dataShape[1] = indicesW * rowCount.value();
3499 }
3500 }
3501
3502 inferredReturnShapes.push_back(ShapedTypeComponents(dataShape));
3503 if (values.size() == 1)
3504 return success();
3505
3506 SmallVector<int64_t> scaleShape = dataShape;
3507 const uint32_t blockSize =
3508 BlockSizeAttr::getBlockSizeValue(adaptor.getBlockSize());
3509 if (ShapedType::isStatic(dataShape[2]))
3510 scaleShape[2] = dataShape[2] / blockSize;
3511
3512 inferredReturnShapes.push_back(ShapedTypeComponents(scaleShape));
3513 return success();
3514}
3515
3516LogicalResult tosa::GatherOp::verify() {
3517 if (verifySameElementTypes(*this, /* inType = */ getValues().getType(),
3518 /* outType = */ getOutput().getType())
3519 .failed()) {
3520 return failure();
3521 }
3522
3523 const ShapeAdaptor valuesShape(getValues().getType());
3524 const ShapeAdaptor indicesShape(getIndices().getType());
3525 const ShapeAdaptor outputShape(getOutput().getType());
3526
3527 int64_t n = ShapedType::kDynamic;
3528 int64_t w = ShapedType::kDynamic;
3529 int64_t c = ShapedType::kDynamic;
3530
3531 if (valuesShape.hasRank()) {
3532 n = valuesShape.getDimSize(0);
3533 c = valuesShape.getDimSize(2);
3534 }
3535 if (indicesShape.hasRank()) {
3536 const int64_t indicesN = indicesShape.getDimSize(0);
3537 w = indicesShape.getDimSize(1);
3538 if (n == ShapedType::kDynamic)
3539 n = indicesN;
3540 else if (indicesN != ShapedType::kDynamic && n != indicesN)
3541 return emitOpError() << "requires indices dimension 0 to have size " << n
3542 << ", got " << indicesN;
3543 }
3544 if (outputShape.hasRank()) {
3545 const int64_t outputN = outputShape.getDimSize(0);
3546 const int64_t outputW = outputShape.getDimSize(1);
3547 const int64_t outputC = outputShape.getDimSize(2);
3548 if (n != ShapedType::kDynamic && outputN != ShapedType::kDynamic &&
3549 n != outputN)
3550 return emitOpError() << "requires output dimension 0 to have size " << n
3551 << ", got " << outputN;
3552
3553 if (w != ShapedType::kDynamic && outputW != ShapedType::kDynamic &&
3554 w != outputW)
3555 return emitOpError() << "requires output dimension 1 to have size " << w
3556 << ", got " << outputW;
3557 if (c != ShapedType::kDynamic && outputC != ShapedType::kDynamic &&
3558 c != outputC)
3559 return emitOpError() << "requires output dimension 2 to have size " << c
3560 << ", got " << outputC;
3561 }
3562 return success();
3563}
3564
3565LogicalResult tosa::RowGatherOp::verify() {
3566 if (failed(verifySameElementTypes(*this, /* inType = */ getValues().getType(),
3567 /* outType = */ getOutput().getType())))
3568 return failure();
3569
3570 const FailureOr<int32_t> maybeRowCount =
3572 if (succeeded(maybeRowCount) && maybeRowCount.value() <= 0)
3573 return emitOpError() << "requires row_count to be > 0, got "
3574 << maybeRowCount.value();
3575
3576 int64_t n = ShapedType::kDynamic;
3577 int64_t c = ShapedType::kDynamic;
3578 int64_t w = ShapedType::kDynamic;
3579
3580 const ShapeAdaptor valuesShape(getValues().getType());
3581 if (valuesShape.hasRank()) {
3582 n = valuesShape.getDimSize(0);
3583 c = valuesShape.getDimSize(2);
3584 }
3585
3586 const ShapeAdaptor indicesShape(getIndices().getType());
3587 if (indicesShape.hasRank()) {
3588 if (failed(tryUpdateDimOrFailure(*this, n, indicesShape.getDimSize(0),
3589 "indices", "batch")))
3590 return failure();
3591 w = indicesShape.getDimSize(1);
3592 }
3593
3594 const ShapeAdaptor outputShape(getOutput().getType());
3595 if (outputShape.hasRank()) {
3596 if (failed(tryUpdateDimOrFailure(*this, n, outputShape.getDimSize(0),
3597 "output", "batch")) ||
3598 failed(tryUpdateDimOrFailure(*this, c, outputShape.getDimSize(2),
3599 "output", "channels")))
3600 return failure();
3601
3602 if (succeeded(maybeRowCount) && maybeRowCount.value() > 0 &&
3603 ShapedType::isStatic(w)) {
3604 const int64_t expectedOutputRows = w * maybeRowCount.value();
3605 if (ShapedType::isStatic(outputShape.getDimSize(1)) &&
3606 outputShape.getDimSize(1) != expectedOutputRows)
3607 return emitOpError()
3608 << "requires output dimension to be equal to "
3609 "indices[1]*row_count ("
3610 << expectedOutputRows << "), got " << outputShape.getDimSize(1);
3611 }
3612 }
3613
3614 return success();
3615}
3616
3617LogicalResult tosa::RowGatherBlockScaledOp::verify() {
3618 const OperandRange values = getValues();
3619 const ResultRange output = getOutput();
3620 if (values.empty() || values.size() > 2)
3621 return emitOpError()
3622 << "expects values tensor list length to be 1 or 2, got "
3623 << values.size();
3624 if (output.size() != values.size())
3625 return emitOpError()
3626 << "expects output tensor list length to match values tensor list "
3627 "length, got "
3628 << output.size() << " results for " << values.size()
3629 << " input tensors";
3630
3631 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(getBlockSize());
3632 if (values.size() == 1 && blockSize != 1)
3633 return emitOpError()
3634 << "requires block_size to be BLOCK_SIZE_1 when values tensor list "
3635 "length is 1";
3636 if (values.size() == 2 && blockSize == 1)
3637 return emitOpError()
3638 << "requires block_size to not be BLOCK_SIZE_1 when values tensor "
3639 "list length is 2";
3640
3641 if (failed(verifySameElementTypes(*this, values[0].getType(),
3642 output[0].getType(), "values[0]",
3643 "output[0]")))
3644 return failure();
3645 if (values.size() == 2 && failed(verifySameElementTypes(
3646 *this, values[1].getType(), output[1].getType(),
3647 "values[1]", "output[1]")))
3648 return failure();
3649
3650 if (auto rowCount = getConstantScalarIntValue<int32_t>(getRowCount());
3651 succeeded(rowCount) && rowCount.value() <= 0)
3652 return emitOpError() << "requires row_count to be > 0, got "
3653 << rowCount.value();
3654
3655 int64_t n = ShapedType::kDynamic;
3656 int64_t k = ShapedType::kDynamic;
3657 int64_t c = ShapedType::kDynamic;
3658 int64_t w = ShapedType::kDynamic;
3659 int64_t multiplesOfC = ShapedType::kDynamic;
3660
3661 const ShapeAdaptor valuesDataShape(values[0].getType());
3662 if (valuesDataShape.hasRank()) {
3663 n = valuesDataShape.getDimSize(0);
3664 k = valuesDataShape.getDimSize(1);
3665 c = valuesDataShape.getDimSize(2);
3666 }
3667
3668 if (ShapedType::isStatic(c) && c % blockSize != 0)
3669 return emitOpError() << "expects channels of values[0] (" << c
3670 << ") to be divisible by block_size (" << blockSize
3671 << ")";
3672
3673 const ShapeAdaptor indicesShape(getIndices().getType());
3674 if (indicesShape.hasRank()) {
3675 if (failed(tryUpdateDimOrFailure(*this, n, indicesShape.getDimSize(0),
3676 "indices", "batch")))
3677 return failure();
3678 w = indicesShape.getDimSize(1);
3679 }
3680
3681 const ShapeAdaptor outputDataShape(output[0].getType());
3682 if (outputDataShape.hasRank()) {
3683 if (failed(tryUpdateDimOrFailure(*this, n, outputDataShape.getDimSize(0),
3684 "output[0]", "batch")) ||
3685 failed(tryUpdateDimOrFailure(*this, c, outputDataShape.getDimSize(2),
3686 "output[0]", "channels")))
3687 return failure();
3688
3689 if (auto rowCount = getConstantScalarIntValue<int32_t>(getRowCount());
3690 succeeded(rowCount) && rowCount.value() > 0 &&
3691 ShapedType::isStatic(w)) {
3692 const int64_t expectedOutputRows = w * rowCount.value();
3693 if (ShapedType::isStatic(outputDataShape.getDimSize(1)) &&
3694 outputDataShape.getDimSize(1) != expectedOutputRows)
3695 return emitOpError() << "requires output[0] dimension 1 to have size "
3696 << expectedOutputRows << ", got "
3697 << outputDataShape.getDimSize(1);
3698 }
3699 }
3700
3701 if (values.size() == 2) {
3702 const ShapeAdaptor valuesScaleShape(values[1].getType());
3703 if (valuesScaleShape.hasRank()) {
3704 if (failed(tryUpdateDimOrFailure(*this, n, valuesScaleShape.getDimSize(0),
3705 "values[1]", "batch")) ||
3706 failed(tryUpdateDimOrFailure(*this, k, valuesScaleShape.getDimSize(1),
3707 "values[1]", "rows")))
3708 return failure();
3709 multiplesOfC = valuesScaleShape.getDimSize(2);
3710 }
3711
3712 const ShapeAdaptor outputScaleShape(output[1].getType());
3713 if (outputScaleShape.hasRank()) {
3714 if (failed(tryUpdateDimOrFailure(*this, n, outputScaleShape.getDimSize(0),
3715 "output[1]", "batch")))
3716 return failure();
3717
3718 if (auto rowCount = getConstantScalarIntValue<int32_t>(getRowCount());
3719 succeeded(rowCount) && rowCount.value() > 0 &&
3720 ShapedType::isStatic(w)) {
3721 const int64_t expectedOutputRows = w * rowCount.value();
3722 if (ShapedType::isStatic(outputScaleShape.getDimSize(1)) &&
3723 outputScaleShape.getDimSize(1) != expectedOutputRows)
3724 return emitOpError() << "requires output[1] dimension 1 to have size "
3725 << expectedOutputRows << ", got "
3726 << outputScaleShape.getDimSize(1);
3727 }
3728
3729 if (ShapedType::isDynamic(multiplesOfC))
3730 multiplesOfC = outputScaleShape.getDimSize(2);
3731 else if (ShapedType::isStatic(outputScaleShape.getDimSize(2)) &&
3732 multiplesOfC != outputScaleShape.getDimSize(2))
3733 return emitOpError()
3734 << "expected channels of output[1] to match size "
3735 << multiplesOfC << ", got " << outputScaleShape.getDimSize(2);
3736 }
3737
3738 if (ShapedType::isStatic(c) && ShapedType::isStatic(multiplesOfC) &&
3739 multiplesOfC != c / blockSize)
3740 return emitOpError()
3741 << "expects channels of scale tensors to equal C/block_size (" << c
3742 << "/" << blockSize << "), got " << multiplesOfC;
3743 }
3744
3745 return success();
3746}
3747
3748LogicalResult tosa::ResizeOp::inferReturnTypeComponents(
3749 MLIRContext *context, ::std::optional<Location> location,
3750 ResizeOp::Adaptor adaptor,
3751 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3752 llvm::SmallVector<int64_t, 4> outputShape;
3753 outputShape.resize(4, ShapedType::kDynamic);
3754
3755 ShapeAdaptor inputShape(adaptor.getInput().getType());
3756 if (!inputShape.hasRank())
3757 return failure();
3758
3759 outputShape[0] = inputShape.getDimSize(0);
3760 outputShape[3] = inputShape.getDimSize(3);
3761 int64_t inputHeight = inputShape.getDimSize(1);
3762 int64_t inputWidth = inputShape.getDimSize(2);
3763
3764 if ((inputHeight == ShapedType::kDynamic) ||
3765 (inputWidth == ShapedType::kDynamic))
3766 return failure();
3767
3768 SmallVector<int64_t> scaleInt, offsetInt, borderInt;
3769 if (!tosa::getConstShapeValues(adaptor.getScale().getDefiningOp(),
3770 scaleInt) ||
3771 !tosa::getConstShapeValues(adaptor.getOffset().getDefiningOp(),
3772 offsetInt) ||
3773 !tosa::getConstShapeValues(adaptor.getBorder().getDefiningOp(),
3774 borderInt)) {
3775 return failure();
3776 }
3777
3778 // Compute the output shape based on attributes: scale, offset, and border.
3779 const int64_t outputHeight =
3780 (((inputHeight - 1) * scaleInt[0] - offsetInt[0] + borderInt[0]) /
3781 scaleInt[1]) +
3782 1;
3783
3784 const int64_t outputWidth =
3785 (((inputWidth - 1) * scaleInt[2] - offsetInt[1] + borderInt[1]) /
3786 scaleInt[3]) +
3787 1;
3788
3789 if (outputHeight < 0 || outputWidth < 0) {
3790 return emitOptionalError(
3791 location,
3792 "calculated output height and width must be non-negative, "
3793 "got height = ",
3794 outputHeight, ", width = ", outputWidth);
3795 }
3796
3797 outputShape[1] = outputHeight;
3798 outputShape[2] = outputWidth;
3799 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3800 return success();
3801}
3802
3803LogicalResult tosa::ResizeOp::verify() {
3804 const Value input = getInput();
3805 const Value output = getOutput();
3806 const Type inputElementType = getElementTypeOrSelf(input.getType());
3807
3808 if (isa<BlockScaledType>(inputElementType) &&
3809 getMode() != ResizeMode::NEAREST_NEIGHBOR)
3810 return emitOpError("requires NEAREST_NEIGHBOR mode for block scaled input");
3811
3812 const RankedTensorType inputType =
3813 llvm::dyn_cast<RankedTensorType>(input.getType());
3814 const RankedTensorType outputType =
3815 llvm::dyn_cast<RankedTensorType>(output.getType());
3816
3817 SmallVector<int64_t> scaleValues;
3818 SmallVector<int64_t> offsetValues;
3819 SmallVector<int64_t> borderValues;
3820 if (!tosa::getConstShapeValues(getScale().getDefiningOp(), scaleValues) ||
3821 !tosa::getConstShapeValues(getOffset().getDefiningOp(), offsetValues) ||
3822 !tosa::getConstShapeValues(getBorder().getDefiningOp(), borderValues)) {
3823 // Skip following checks if shape is not constant
3824 return success();
3825 }
3826
3827 if (llvm::any_of(scaleValues, [](int64_t s) { return s <= 0; }))
3828 return emitOpError("expect all scale values to be > 0, got ")
3829 << scaleValues;
3830
3831 const int64_t scaleYN = scaleValues[0];
3832 const int64_t scaleYD = scaleValues[1];
3833 const int64_t scaleXN = scaleValues[2];
3834 const int64_t scaleXD = scaleValues[3];
3835
3836 const int64_t offsetY = offsetValues[0];
3837 const int64_t offsetX = offsetValues[1];
3838
3839 const int64_t borderY = borderValues[0];
3840 const int64_t borderX = borderValues[1];
3841
3842 if (!inputType)
3843 return success();
3844 if (!outputType)
3845 return success();
3846
3847 const int64_t oh = outputType.getDimSize(1);
3848 const int64_t ow = outputType.getDimSize(2);
3849 const int64_t ih = inputType.getDimSize(1);
3850 const int64_t iw = inputType.getDimSize(2);
3851
3852 // Don't check with input height that could be broadcast (ih != 1)
3853 // since Linalg, a consumer of TOSA, expects broadcasting support
3854 // in resize to be available. Taking the cautious approach for now,
3855 // we can consider removing support for broadcasting later.
3856 if (ih != ShapedType::kDynamic && ih != 1) {
3857 const std::optional<int64_t> calculatedOutHeightMinusOne =
3858 idivCheck((ih - 1) * scaleYN - offsetY + borderY, scaleYD);
3859 if (!calculatedOutHeightMinusOne.has_value())
3860 return emitOpError("expected (input_height - 1) * scale_y_n - offset_y + "
3861 "border_y ")
3862 << "to be wholly divisible by scale_y_d, got ((" << ih
3863 << " - 1) * " << scaleYN << " - " << offsetY << " + " << borderY
3864 << ") / " << scaleYD;
3865 const int64_t calculatedOutHeight = calculatedOutHeightMinusOne.value() + 1;
3866 if (oh != ShapedType::kDynamic && calculatedOutHeight != oh)
3867 return emitOpError("calculated output height did not match expected: ")
3868 << "calculated=" << calculatedOutHeight << ", expected=" << oh;
3869 }
3870
3871 // Don't check with input width that could be broadcast (iw != 1)
3872 // since Linalg, a consumer of TOSA, expects broadcasting support
3873 // in resize to be available. Taking the cautious approach for now,
3874 // we can consider removing support for broadcasting later.
3875 if (iw != ShapedType::kDynamic && iw != 1) {
3876 const int64_t scaledInWidth = (iw - 1) * scaleXN - offsetX + borderX;
3877 const std::optional<int64_t> calculatedOutWidthMinusOne =
3878 idivCheck(scaledInWidth, scaleXD);
3879 if (!calculatedOutWidthMinusOne.has_value())
3880 return emitOpError("expected (input_width - 1) * scale_x_n - offset_x + "
3881 "border_x ")
3882 << "to be wholly divisible by scale_x_d, got ((" << iw
3883 << " - 1) * " << scaleXN << " - " << offsetX << " + " << borderX
3884 << ") / " << scaleXD;
3885 const int64_t calculatedOutWidth = calculatedOutWidthMinusOne.value() + 1;
3886 if (ow != ShapedType::kDynamic && calculatedOutWidth != ow)
3887 return emitOpError("calculated output width did not match expected: ")
3888 << "calculated=" << calculatedOutWidth << ", expected=" << ow;
3889 }
3890
3891 return success();
3892}
3893
3894LogicalResult tosa::ScatterOp::inferReturnTypeComponents(
3895 MLIRContext *context, ::std::optional<Location> location,
3896 ScatterOp::Adaptor adaptor,
3897 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3898 llvm::SmallVector<int64_t> outputShape;
3899 outputShape.resize(3, ShapedType::kDynamic);
3900
3901 ShapeAdaptor valuesInShape(adaptor.getValuesIn().getType());
3902 if (valuesInShape.hasRank()) {
3903 outputShape[0] = valuesInShape.getDimSize(0);
3904 outputShape[1] = valuesInShape.getDimSize(1);
3905 outputShape[2] = valuesInShape.getDimSize(2);
3906 }
3907
3908 ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3909 if (indicesShape.hasRank()) {
3910 if (outputShape[0] == ShapedType::kDynamic)
3911 outputShape[0] = indicesShape.getDimSize(0);
3912 }
3913
3914 ShapeAdaptor inputShape(adaptor.getInput().getType());
3915 if (inputShape.hasRank()) {
3916 if (outputShape[0] == ShapedType::kDynamic)
3917 outputShape[0] = inputShape.getDimSize(0);
3918 if (outputShape[2] == ShapedType::kDynamic)
3919 outputShape[2] = inputShape.getDimSize(2);
3920 }
3921
3922 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3923 return success();
3924}
3925
3926LogicalResult tosa::ScatterOp::verify() {
3927 if (verifySameElementTypes(*this, /* inType = */ getValuesIn().getType(),
3928 /* outType = */ getValuesOut().getType())
3929 .failed() ||
3930 verifySameElementTypes(*this, /* inType = */ getInput().getType(),
3931 /* outType = */ getValuesOut().getType())
3932 .failed()) {
3933 return failure();
3934 }
3935
3936 const ShapeAdaptor valuesInShape(getValuesIn().getType());
3937 const ShapeAdaptor indicesShape(getIndices().getType());
3938 const ShapeAdaptor inputShape(getInput().getType());
3939 const ShapeAdaptor outputShape(getValuesOut().getType());
3940
3941 int64_t n = ShapedType::kDynamic;
3942 int64_t k = ShapedType::kDynamic;
3943 int64_t w = ShapedType::kDynamic;
3944 int64_t c = ShapedType::kDynamic;
3945 if (valuesInShape.hasRank()) {
3946 n = valuesInShape.getDimSize(0);
3947 k = valuesInShape.getDimSize(1);
3948 c = valuesInShape.getDimSize(2);
3949 }
3950 if (indicesShape.hasRank()) {
3951 const int64_t indicesN = indicesShape.getDimSize(0);
3952 w = indicesShape.getDimSize(1);
3953 if (n == ShapedType::kDynamic)
3954 n = indicesN;
3955 else if (indicesN != ShapedType::kDynamic && n != indicesN)
3956 return emitOpError() << "requires indices dimension 0 to have size " << n
3957 << ", got " << indicesN;
3958 }
3959 if (inputShape.hasRank()) {
3960 const int64_t inputN = inputShape.getDimSize(0);
3961 const int64_t inputW = inputShape.getDimSize(1);
3962 const int64_t inputC = inputShape.getDimSize(2);
3963 if (n == ShapedType::kDynamic)
3964 n = inputN;
3965 else if (inputN != ShapedType::kDynamic && n != inputN)
3966 return emitOpError() << "requires input dimension 0 to have size " << n
3967 << ", got " << inputN;
3968 if (w == ShapedType::kDynamic)
3969 w = inputW;
3970 else if (inputW != ShapedType::kDynamic && w != inputW)
3971 return emitOpError() << "requires input dimension 1 to have size " << w
3972 << ", got " << inputW;
3973
3974 if (c == ShapedType::kDynamic)
3975 c = inputC;
3976 else if (inputC != ShapedType::kDynamic && c != inputC)
3977 return emitOpError() << "requires input dimension 2 to have size " << c
3978 << ", got " << inputC;
3979 }
3980 if (outputShape.hasRank()) {
3981 const int64_t outputN = outputShape.getDimSize(0);
3982 const int64_t outputK = outputShape.getDimSize(1);
3983 const int64_t outputC = outputShape.getDimSize(2);
3984 if (n != ShapedType::kDynamic && outputN != ShapedType::kDynamic &&
3985 n != outputN)
3986 return emitOpError() << "requires values_out dimension 0 to have size "
3987 << n << ", got " << outputN;
3988 if (k == ShapedType::kDynamic)
3989 k = outputK;
3990 else if (outputK != ShapedType::kDynamic && k != outputK)
3991 return emitOpError() << "requires values_out dimension 1 to have size "
3992 << k << ", got " << outputK;
3993 if (c != ShapedType::kDynamic && outputC != ShapedType::kDynamic &&
3994 c != outputC)
3995 return emitOpError() << "requires values_out dimension 2 to have size "
3996 << c << ", got " << outputC;
3997 }
3998 if (k != ShapedType::kDynamic && w != ShapedType::kDynamic && !(k >= w))
3999 return emitOpError() << "requires dimensions K >= W, got K=" << k
4000 << " and W=" << w;
4001
4002 return success();
4003}
4004
4005static LogicalResult ReduceInferReturnTypes(
4006 ShapeAdaptor operandShape, Type inputType, IntegerAttr axis,
4007 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4008 int64_t axisVal = axis.getValue().getSExtValue();
4009 if (!operandShape.hasRank() || operandShape.getRank() <= axisVal) {
4010 inferredReturnShapes.push_back(ShapedTypeComponents(inputType));
4011 return success();
4012 }
4013
4014 SmallVector<int64_t> outputShape;
4015 operandShape.getDims(outputShape);
4016 outputShape[axisVal] = 1;
4017 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape, inputType));
4018 return success();
4019}
4020
4021#define COMPATIBLE_RETURN_TYPES(OP) \
4022 bool OP::isCompatibleReturnTypes(TypeRange l, TypeRange r) { \
4023 if (l.size() != r.size() || l.size() != 1) \
4024 return false; \
4025 if (getElementTypeOrSelf(l[0]) != getElementTypeOrSelf(r[0])) \
4026 return false; \
4027 return succeeded(verifyCompatibleShape(l[0], r[0])); \
4028 }
4029
4030#define REDUCE_SHAPE_INFER(OP) \
4031 LogicalResult OP::inferReturnTypeComponents( \
4032 MLIRContext *context, ::std::optional<Location> location, \
4033 OP::Adaptor adaptor, \
4034 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) { \
4035 Type inputType = \
4036 llvm::cast<TensorType>(adaptor.getInput().getType()).getElementType(); \
4037 ShapeAdaptor inputShape(adaptor.getInput().getType()); \
4038 const Properties &prop = adaptor.getProperties(); \
4039 return ReduceInferReturnTypes(inputShape, inputType, prop.axis, \
4040 inferredReturnShapes); \
4041 } \
4042 COMPATIBLE_RETURN_TYPES(OP)
4043
4044REDUCE_SHAPE_INFER(tosa::ReduceAllOp)
4045REDUCE_SHAPE_INFER(tosa::ReduceAnyOp)
4046REDUCE_SHAPE_INFER(tosa::ReduceMaxOp)
4047REDUCE_SHAPE_INFER(tosa::ReduceMinOp)
4048REDUCE_SHAPE_INFER(tosa::ReduceProductOp)
4049REDUCE_SHAPE_INFER(tosa::ReduceSumOp)
4050#undef REDUCE_SHAPE_INFER
4051COMPATIBLE_RETURN_TYPES(tosa::ConcatOp)
4052#undef COMPATIBLE_RETURN_TYPES
4053
4054template <typename T>
4055static LogicalResult verifyReduceOp(T op) {
4056 // All TOSA reduce Ops have input, output and axis.
4057 TensorType inputType = op.getInput().getType();
4058 TensorType outputType = op.getOutput().getType();
4059 int32_t reduceAxis = op.getAxis();
4060
4061 if (reduceAxis < 0) {
4062 op.emitOpError("reduce axis must not be negative");
4063 return failure();
4064 }
4065 if (inputType.hasRank()) {
4066 int64_t inputRank = inputType.getRank();
4067 // We allow for a special case where the input/output shape has rank 0 and
4068 // axis is also 0.
4069 if (reduceAxis >= inputRank && (reduceAxis != 0 || inputRank != 0)) {
4070 op.emitOpError("expect input tensor rank (")
4071 << inputRank << ") to be larger than reduce axis (" << reduceAxis
4072 << ")";
4073 return failure();
4074 }
4075 }
4076 if (outputType.hasRank()) {
4077 int64_t outputRank = outputType.getRank();
4078 if (inputType.hasRank() && outputRank != inputType.getRank()) {
4079 op.emitOpError(
4080 "expect output tensor rank to be equal to input tensor rank");
4081 return failure();
4082 }
4083 if (reduceAxis >= outputRank && (reduceAxis != 0 || outputRank != 0)) {
4084 op.emitOpError("expect output tensor rank (")
4085 << outputRank << ") to be larger than reduce axis (" << reduceAxis
4086 << ")";
4087 return failure();
4088 }
4089 // We can only verify the reduced dimension size to be 1 if this is not
4090 // the special case of output rank == 0.
4091 if (outputRank != 0) {
4092 auto outputShape = outputType.getShape();
4093 if (!outputType.isDynamicDim(reduceAxis) &&
4094 outputShape[reduceAxis] != 1) {
4095 op.emitOpError("expect reduced dimension size to be 1, got ")
4096 << outputShape[reduceAxis];
4097 return failure();
4098 }
4099 }
4100 }
4101 return success();
4102}
4103
4104LogicalResult tosa::ReduceAllOp::verify() { return verifyReduceOp(*this); }
4105LogicalResult tosa::ReduceAnyOp::verify() { return verifyReduceOp(*this); }
4106LogicalResult tosa::ReduceMaxOp::verify() { return verifyReduceOp(*this); }
4107LogicalResult tosa::ReduceMinOp::verify() { return verifyReduceOp(*this); }
4108LogicalResult tosa::ReduceProductOp::verify() { return verifyReduceOp(*this); }
4109LogicalResult tosa::ReduceSumOp::verify() { return verifyReduceOp(*this); }
4110
4111static LogicalResult NAryInferReturnTypes(
4112 const ValueShapeRange &operands,
4113 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4115 if (resolveBroadcastShape(operands, outShape).failed()) {
4116 inferredReturnShapes.push_back(ShapedTypeComponents());
4117 } else {
4118 inferredReturnShapes.push_back(ShapedTypeComponents(outShape));
4119 }
4120 return success();
4121}
4122
4123#define NARY_SHAPE_INFER(OP) \
4124 LogicalResult OP::inferReturnTypeComponents( \
4125 MLIRContext *context, ::std::optional<Location> location, \
4126 ValueShapeRange operands, DictionaryAttr attributes, \
4127 PropertyRef properties, RegionRange regions, \
4128 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) { \
4129 return NAryInferReturnTypes(operands, inferredReturnShapes); \
4130 }
4131
4132NARY_SHAPE_INFER(tosa::AbsOp)
4133NARY_SHAPE_INFER(tosa::AddOp)
4134NARY_SHAPE_INFER(tosa::ArithmeticRightShiftOp)
4135NARY_SHAPE_INFER(tosa::BitwiseAndOp)
4136NARY_SHAPE_INFER(tosa::BitwiseOrOp)
4137NARY_SHAPE_INFER(tosa::BitwiseXorOp)
4138NARY_SHAPE_INFER(tosa::BitwiseNotOp)
4139NARY_SHAPE_INFER(tosa::CastOp)
4140NARY_SHAPE_INFER(tosa::CeilOp)
4141NARY_SHAPE_INFER(tosa::ClampOp)
4142NARY_SHAPE_INFER(tosa::ClzOp)
4143NARY_SHAPE_INFER(tosa::CosOp)
4144NARY_SHAPE_INFER(tosa::ExpOp)
4145NARY_SHAPE_INFER(tosa::FloorOp)
4146NARY_SHAPE_INFER(tosa::GreaterEqualOp)
4147NARY_SHAPE_INFER(tosa::GreaterOp)
4148NARY_SHAPE_INFER(tosa::IdentityOp)
4149NARY_SHAPE_INFER(tosa::IntDivOp)
4150NARY_SHAPE_INFER(tosa::LogOp)
4151NARY_SHAPE_INFER(tosa::LogicalAndOp)
4152NARY_SHAPE_INFER(tosa::LogicalLeftShiftOp)
4153NARY_SHAPE_INFER(tosa::LogicalNotOp)
4154NARY_SHAPE_INFER(tosa::LogicalOrOp)
4155NARY_SHAPE_INFER(tosa::LogicalRightShiftOp)
4156NARY_SHAPE_INFER(tosa::LogicalXorOp)
4157NARY_SHAPE_INFER(tosa::MaximumOp)
4158NARY_SHAPE_INFER(tosa::MinimumOp)
4159NARY_SHAPE_INFER(tosa::PowOp)
4160NARY_SHAPE_INFER(tosa::ReciprocalOp)
4161NARY_SHAPE_INFER(tosa::ReverseOp)
4162NARY_SHAPE_INFER(tosa::RsqrtOp)
4163NARY_SHAPE_INFER(tosa::SinOp)
4164NARY_SHAPE_INFER(tosa::SelectOp)
4165NARY_SHAPE_INFER(tosa::SubOp)
4166NARY_SHAPE_INFER(tosa::TanhOp)
4167NARY_SHAPE_INFER(tosa::ErfOp)
4168NARY_SHAPE_INFER(tosa::SigmoidOp)
4169#undef PRED_SHAPE_INFER
4170
4171LogicalResult tosa::NegateOp::inferReturnTypeComponents(
4172 MLIRContext *context, ::std::optional<Location> location,
4173 NegateOp::Adaptor adaptor,
4174 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4175 ShapeAdaptor inputShape(adaptor.getInput1().getType());
4176 inferredReturnShapes.push_back(ShapedTypeComponents(inputShape));
4177 return success();
4178}
4179
4180LogicalResult tosa::NegateOp::verify() {
4181 // Verify same element type
4182 const Type input1Type = getInput1().getType();
4183 const Type outputType = getOutput().getType();
4184 if (verifySameElementTypes(*this, input1Type, outputType).failed())
4185 return failure();
4186
4187 // Verify same shape
4188 const SmallVector<Type, 2> types = {input1Type, outputType};
4189 if (failed(verifyCompatibleShapes(types)))
4190 return emitOpError() << "requires the same shape for input1 and output";
4191
4192 const Type input1EType = getStorageElementTypeOrSelf(getInput1().getType());
4193 const Type input1ZpEType =
4194 getStorageElementTypeOrSelf(getInput1Zp().getType());
4195 if (input1EType != input1ZpEType) {
4196 return emitOpError("expect both input1 and its zero point are the same "
4197 "element type, got ")
4198 << input1EType << " and " << input1ZpEType;
4199 }
4200 const Type outputEType = getStorageElementTypeOrSelf(getOutput().getType());
4201 const Type outputZpEType =
4202 getStorageElementTypeOrSelf(getOutputZp().getType());
4203 if (outputEType != outputZpEType) {
4204 return emitOpError("expect both output and its zero point are the same "
4205 "element type, got ")
4206 << outputEType << " and " << outputZpEType;
4207 }
4208
4209 FailureOr<int64_t> maybeIZp = getInput1ZeroPoint();
4210 if (succeeded(maybeIZp) && verifyInput1ZeroPoint(*maybeIZp).failed())
4211 return failure();
4212
4213 FailureOr<int64_t> maybeOZp = getOutputZeroPoint();
4214 if (succeeded(maybeOZp) && verifyOutputZeroPoint(*maybeOZp).failed())
4215 return failure();
4216
4217 return success();
4218}
4219
4220static LogicalResult poolingInferReturnTypes(
4221 ShapeAdaptor inputShape, ArrayRef<int64_t> kernel, ArrayRef<int64_t> stride,
4223 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4224 llvm::SmallVector<int64_t> outputShape;
4225 outputShape.resize(4, ShapedType::kDynamic);
4226
4227 // We only know the rank if the input type is unranked.
4228 if (!inputShape.hasRank()) {
4229 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4230 return success();
4231 }
4232
4233 // Batch and number of channels are identical for pooling layer.
4234 outputShape[0] = inputShape.getDimSize(0);
4235 outputShape[3] = inputShape.getDimSize(3);
4236
4237 int64_t height = inputShape.getDimSize(1);
4238 int64_t width = inputShape.getDimSize(2);
4239
4240 if (ShapedType::isStatic(height)) {
4241 int64_t padded = height + pad[0] + pad[1] - kernel[0];
4242 outputShape[1] = padded / stride[0] + 1;
4243 }
4244
4245 if (ShapedType::isStatic(width)) {
4246 int64_t padded = width + pad[2] + pad[3] - kernel[1];
4247 outputShape[2] = padded / stride[1] + 1;
4248 }
4249
4250 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4251 return success();
4252}
4253
4254template <typename AdaptorT>
4256
4258protected:
4259 static void updateIfDynamic(int64_t &current, int64_t candidate) {
4260 if (ShapedType::isDynamic(current))
4261 current = candidate;
4262 }
4263};
4264
4265template <>
4266class ConvInferShapeAdaptor<Conv2DOp::Adaptor>
4267 : public ConvInferShapeAdaptorBase {
4268public:
4269 explicit ConvInferShapeAdaptor(Conv2DOp::Adaptor adaptor)
4270 : adaptor(adaptor) {}
4271
4273 SmallVectorImpl<int64_t> &inputSpatial) {
4274 const ShapeAdaptor inputShape(adaptor.getInput().getType());
4275 if (!inputShape.hasRank())
4276 return;
4277
4278 const int64_t outputBatch = inputShape.getDimSize(0);
4279 const int64_t inputHeight = inputShape.getDimSize(1);
4280 const int64_t inputWidth = inputShape.getDimSize(2);
4281
4282 outputShape[0] = outputBatch;
4283 inputSpatial[0] = inputHeight;
4284 inputSpatial[1] = inputWidth;
4285 }
4286
4288 SmallVectorImpl<int64_t> &weightSpatial) {
4289 const ShapeAdaptor weightShape(adaptor.getWeight().getType());
4290 if (!weightShape.hasRank())
4291 return;
4292
4293 const int64_t outputChannels = weightShape.getDimSize(0);
4294 const int64_t kernelHeight = weightShape.getDimSize(1);
4295 const int64_t kernelWidth = weightShape.getDimSize(2);
4296
4297 outputShape[3] = outputChannels;
4298 weightSpatial[0] = kernelHeight;
4299 weightSpatial[1] = kernelWidth;
4300 }
4301
4302 int64_t getNumSpatialDims() const { return 2; }
4303 int64_t getOutputRank() const { return 4; }
4304
4306 SmallVector<int64_t> &strideValues,
4307 SmallVector<int64_t> &dilationValues) {
4308 padValues.assign(adaptor.getPad().begin(), adaptor.getPad().end());
4309 strideValues.assign(adaptor.getStride().begin(), adaptor.getStride().end());
4310 dilationValues.assign(adaptor.getDilation().begin(),
4311 adaptor.getDilation().end());
4312 return success();
4313 }
4314
4315private:
4316 Conv2DOp::Adaptor adaptor;
4317};
4318
4319template <>
4320class ConvInferShapeAdaptor<Conv2DBlockScaledOp::Adaptor>
4321 : public ConvInferShapeAdaptorBase {
4322public:
4323 explicit ConvInferShapeAdaptor(Conv2DBlockScaledOp::Adaptor adaptor)
4324 : adaptor(adaptor) {}
4325
4327 SmallVectorImpl<int64_t> &inputSpatial) {
4328 const ShapeAdaptor inputDataShape(adaptor.getInputData().getType());
4329 if (inputDataShape.hasRank()) {
4330 const int64_t outputBatch = inputDataShape.getDimSize(0);
4331 const int64_t inputHeight = inputDataShape.getDimSize(1);
4332 const int64_t inputWidth = inputDataShape.getDimSize(2);
4333
4334 outputShape[0] = outputBatch;
4335 inputSpatial[0] = inputHeight;
4336 inputSpatial[1] = inputWidth;
4337 }
4338
4339 const ShapeAdaptor inputScaleShape(adaptor.getInputScale().getType());
4340 if (!inputScaleShape.hasRank())
4341 return;
4342
4343 const int64_t scaleBatch = inputScaleShape.getDimSize(0);
4344 const int64_t scaleHeight = inputScaleShape.getDimSize(1);
4345 const int64_t scaleWidth = inputScaleShape.getDimSize(2);
4346
4347 updateIfDynamic(outputShape[0], scaleBatch);
4348 updateIfDynamic(inputSpatial[0], scaleHeight);
4349 updateIfDynamic(inputSpatial[1], scaleWidth);
4350 }
4351
4353 SmallVectorImpl<int64_t> &weightSpatial) {
4354 const ShapeAdaptor weightDataShape(adaptor.getWeightData().getType());
4355 if (weightDataShape.hasRank()) {
4356 const int64_t outputChannels = weightDataShape.getDimSize(0);
4357 const int64_t kernelHeight = weightDataShape.getDimSize(1);
4358 const int64_t kernelWidth = weightDataShape.getDimSize(2);
4359
4360 outputShape[3] = outputChannels;
4361 weightSpatial[0] = kernelHeight;
4362 weightSpatial[1] = kernelWidth;
4363 }
4364
4365 const ShapeAdaptor weightScaleShape(adaptor.getWeightScale().getType());
4366 if (!weightScaleShape.hasRank())
4367 return;
4368
4369 const int64_t scaleOutputChannels = weightScaleShape.getDimSize(0);
4370 const int64_t scaleKernelHeight = weightScaleShape.getDimSize(1);
4371 const int64_t scaleKernelWidth = weightScaleShape.getDimSize(2);
4372
4373 updateIfDynamic(outputShape[3], scaleOutputChannels);
4374 updateIfDynamic(weightSpatial[0], scaleKernelHeight);
4375 updateIfDynamic(weightSpatial[1], scaleKernelWidth);
4376 }
4377
4378 int64_t getNumSpatialDims() const { return 2; }
4379 int64_t getOutputRank() const { return 4; }
4380
4382 SmallVector<int64_t> &strideValues,
4383 SmallVector<int64_t> &dilationValues) {
4384 if (!tosa::getConstShapeValues(adaptor.getPad().getDefiningOp(),
4385 padValues) ||
4386 !tosa::getConstShapeValues(adaptor.getStride().getDefiningOp(),
4387 strideValues) ||
4388 !tosa::getConstShapeValues(adaptor.getDilation().getDefiningOp(),
4389 dilationValues))
4390 return failure();
4391 return success();
4392 }
4393
4394private:
4395 Conv2DBlockScaledOp::Adaptor adaptor;
4396};
4397
4398template <>
4399class ConvInferShapeAdaptor<Conv3DOp::Adaptor>
4400 : public ConvInferShapeAdaptorBase {
4401public:
4402 explicit ConvInferShapeAdaptor(Conv3DOp::Adaptor adaptor)
4403 : adaptor(adaptor) {}
4404
4406 SmallVectorImpl<int64_t> &inputSpatial) {
4407 const ShapeAdaptor inputShape(adaptor.getInput().getType());
4408 if (!inputShape.hasRank())
4409 return;
4410
4411 const int64_t outputBatch = inputShape.getDimSize(0);
4412 const int64_t inputDepth = inputShape.getDimSize(1);
4413 const int64_t inputHeight = inputShape.getDimSize(2);
4414 const int64_t inputWidth = inputShape.getDimSize(3);
4415
4416 outputShape[0] = outputBatch;
4417 inputSpatial[0] = inputDepth;
4418 inputSpatial[1] = inputHeight;
4419 inputSpatial[2] = inputWidth;
4420 }
4421
4423 SmallVectorImpl<int64_t> &weightSpatial) {
4424 const ShapeAdaptor weightShape(adaptor.getWeight().getType());
4425 if (!weightShape.hasRank())
4426 return;
4427
4428 const int64_t outputChannels = weightShape.getDimSize(0);
4429 const int64_t kernelDepth = weightShape.getDimSize(1);
4430 const int64_t kernelHeight = weightShape.getDimSize(2);
4431 const int64_t kernelWidth = weightShape.getDimSize(3);
4432
4433 outputShape[4] = outputChannels;
4434 weightSpatial[0] = kernelDepth;
4435 weightSpatial[1] = kernelHeight;
4436 weightSpatial[2] = kernelWidth;
4437 }
4438
4439 int64_t getNumSpatialDims() const { return 3; }
4440 int64_t getOutputRank() const { return 5; }
4441
4443 SmallVector<int64_t> &strideValues,
4444 SmallVector<int64_t> &dilationValues) {
4445 padValues.assign(adaptor.getPad().begin(), adaptor.getPad().end());
4446 strideValues.assign(adaptor.getStride().begin(), adaptor.getStride().end());
4447 dilationValues.assign(adaptor.getDilation().begin(),
4448 adaptor.getDilation().end());
4449 return success();
4450 }
4451
4452private:
4453 Conv3DOp::Adaptor adaptor;
4454};
4455
4456template <typename AdaptorT>
4458 AdaptorT adaptor,
4459 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4460 ConvInferShapeAdaptor<AdaptorT> convShapeAdaptor(adaptor);
4461 llvm::SmallVector<int64_t> outputShape(convShapeAdaptor.getOutputRank(),
4462 ShapedType::kDynamic);
4463 llvm::SmallVector<int64_t> inputSpatial(convShapeAdaptor.getNumSpatialDims(),
4464 ShapedType::kDynamic);
4465 llvm::SmallVector<int64_t> weightSpatial(convShapeAdaptor.getNumSpatialDims(),
4466 ShapedType::kDynamic);
4467
4468 convShapeAdaptor.inferInputShape(outputShape, inputSpatial);
4469 convShapeAdaptor.inferWeightShape(outputShape, weightSpatial);
4470
4471 const ShapeAdaptor biasShape = adaptor.getBias().getType();
4472 if (biasShape.hasRank()) {
4473 const int64_t biasSize = biasShape.getDimSize(0);
4474 if (biasSize != 1) {
4475 const size_t outputChannelDim = convShapeAdaptor.getOutputRank() - 1;
4476 outputShape[outputChannelDim] =
4477 ShapedType::isDynamic(outputShape[outputChannelDim])
4478 ? biasSize
4479 : outputShape[outputChannelDim];
4480 }
4481 }
4482
4483 SmallVector<int64_t> padValues;
4484 SmallVector<int64_t> strideValues;
4485 SmallVector<int64_t> dilationValues;
4486 if (failed(convShapeAdaptor.getSpatialParameters(padValues, strideValues,
4487 dilationValues))) {
4488 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4489 return success();
4490 }
4491
4492 for (int64_t dim = 0; dim < convShapeAdaptor.getNumSpatialDims(); ++dim) {
4493 if (!ShapedType::isStatic(inputSpatial[dim]) ||
4494 !ShapedType::isStatic(weightSpatial[dim]))
4495 continue;
4496 const int64_t inputSize =
4497 inputSpatial[dim] + padValues[2 * dim] + padValues[2 * dim + 1];
4498 const int64_t filterSize =
4499 (weightSpatial[dim] - 1) * dilationValues[dim] + 1;
4500 const int64_t unstridedResult = inputSize - filterSize + 1;
4501 outputShape[dim + 1] = (unstridedResult - 1) / strideValues[dim] + 1;
4502 }
4503
4504 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4505 return success();
4506}
4507
4508LogicalResult Conv2DOp::inferReturnTypeComponents(
4509 MLIRContext *context, ::std::optional<Location> location,
4510 Conv2DOp::Adaptor adaptor,
4511 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4512 return inferConvReturnTypeComponents(adaptor, inferredReturnShapes);
4513}
4514
4515LogicalResult Conv2DOp::verify() {
4516 if (verifyConvOp(*this).failed() || verifyConvOpModes(*this).failed() ||
4517 verifyConvOpErrorIf(*this).failed())
4518 return failure();
4519 return success();
4520}
4521
4522LogicalResult Conv2DBlockScaledOp::inferReturnTypeComponents(
4523 MLIRContext *context, ::std::optional<Location> location,
4524 Conv2DBlockScaledOp::Adaptor adaptor,
4525 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4526 return inferConvReturnTypeComponents(adaptor, inferredReturnShapes);
4527}
4528
4529LogicalResult Conv2DBlockScaledOp::verify() {
4530 if (failed(verifySameElementTypes(*this, getInputData().getType(),
4531 getWeightData().getType(), "input_data",
4532 "weight_data")) ||
4533 failed(verifySameElementTypes(*this, getInputScale().getType(),
4534 getWeightScale().getType(), "input_scale",
4535 "weight_scale")) ||
4536 failed(verifySameElementTypes(*this, getBias().getType(),
4537 getOutput().getType(), "bias", "output")))
4538 return failure();
4539
4540 // Verify input shape compatibility
4541 int64_t N = ShapedType::kDynamic;
4542 int64_t IH = ShapedType::kDynamic;
4543 int64_t IW = ShapedType::kDynamic;
4544 int64_t IC = ShapedType::kDynamic;
4545 int64_t multiplesOfIC = ShapedType::kDynamic;
4546 int64_t OC = ShapedType::kDynamic;
4547 int64_t KH = ShapedType::kDynamic;
4548 int64_t KW = ShapedType::kDynamic;
4549
4550 const ShapeAdaptor inputDataShape(getInputData().getType());
4551 if (inputDataShape.hasRank()) {
4552 N = inputDataShape.getDimSize(0);
4553 IH = inputDataShape.getDimSize(1);
4554 IW = inputDataShape.getDimSize(2);
4555 IC = inputDataShape.getDimSize(3);
4556 }
4557
4558 const ShapeAdaptor inputScaleShape(getInputScale().getType());
4559 if (inputScaleShape.hasRank()) {
4560 if (failed(tryUpdateDimOrFailure(*this, N, inputScaleShape.getDimSize(0),
4561 "input_scale", "batch size")) ||
4562 failed(tryUpdateDimOrFailure(*this, IH, inputScaleShape.getDimSize(1),
4563 "input_scale", "input height")) ||
4564 failed(tryUpdateDimOrFailure(*this, IW, inputScaleShape.getDimSize(2),
4565 "input_scale", "input width")))
4566 return failure();
4567 multiplesOfIC = inputScaleShape.getDimSize(3);
4568 }
4569
4570 const ShapeAdaptor weightDataShape(getWeightData().getType());
4571 if (weightDataShape.hasRank()) {
4572 OC = weightDataShape.getDimSize(0);
4573 KH = weightDataShape.getDimSize(1);
4574 KW = weightDataShape.getDimSize(2);
4575 if (failed(tryUpdateDimOrFailure(*this, IC, weightDataShape.getDimSize(3),
4576 "weight_data", "input channels")))
4577 return failure();
4578 }
4579
4580 const ShapeAdaptor weightScaleShape(getWeightScale().getType());
4581 if (weightScaleShape.hasRank()) {
4582 if (failed(tryUpdateDimOrFailure(*this, OC, weightScaleShape.getDimSize(0),
4583 "weight_scale", "output channels")) ||
4584 failed(tryUpdateDimOrFailure(*this, KH, weightScaleShape.getDimSize(1),
4585 "weight_scale", "kernel height")) ||
4586 failed(tryUpdateDimOrFailure(*this, KW, weightScaleShape.getDimSize(2),
4587 "weight_scale", "kernel width")) ||
4588 failed(tryUpdateDimOrFailure(*this, multiplesOfIC,
4589 weightScaleShape.getDimSize(3),
4590 "weight_scale", "input channel blocks")))
4591 return failure();
4592 }
4593
4594 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(getBlockSize());
4595 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
4596 return emitOpError("expect block size to be 32, got ") << blockSize;
4597 // Verify IC is a multiple of block size
4598 if (ShapedType::isStatic(IC) && IC % blockSize != 0)
4599 return emitOpError("expect IC to be a multiple of block size, got IC=")
4600 << IC << ", block_size=" << blockSize;
4601
4602 // Verify multiplesOfIC is IC / block size
4603 if (ShapedType::isStatic(IC) && ShapedType::isStatic(multiplesOfIC) &&
4604 multiplesOfIC != IC / blockSize)
4605 return emitOpError(
4606 "expect scale operands dimension 2 to equal IC/block_size (")
4607 << IC << "/" << blockSize << ")"
4608 << ", got " << multiplesOfIC;
4609
4610 // Verify pad/stride/dilation values
4611 SmallVector<int64_t> padValues;
4612 if (tosa::getConstShapeValues(getPad().getDefiningOp(), padValues)) {
4613 if (llvm::any_of(padValues, [](int64_t p) { return p < 0; }))
4614 return emitOpError("expect all padding values to be >= 0, got ")
4615 << padValues;
4616 }
4617
4618 SmallVector<int64_t> strideValues;
4619 if (tosa::getConstShapeValues(getStride().getDefiningOp(), strideValues)) {
4620 if (llvm::any_of(strideValues, [](int64_t s) { return s < 1; }))
4621 return emitOpError("expect all stride values to be >= 1, got ")
4622 << strideValues;
4623 }
4624
4625 SmallVector<int64_t> dilationValues;
4626 if (tosa::getConstShapeValues(getDilation().getDefiningOp(),
4627 dilationValues)) {
4628 if (llvm::any_of(dilationValues, [](int64_t d) { return d < 1; }))
4629 return emitOpError("expect all dilation values to be >= 1, got ")
4630 << dilationValues;
4631 }
4632
4633 // Verify output shape compatibility
4634 const ShapeAdaptor outputShape(getOutput().getType());
4635 if (!padValues.empty() && !strideValues.empty() && !dilationValues.empty() &&
4636 outputShape.hasRank()) {
4637 if (failed(verifyConvOutputSize(*this, IH, KH, outputShape.getDimSize(1),
4638 padValues[0], padValues[1], strideValues[0],
4639 dilationValues[0], "height", "y", "top",
4640 "bottom")) ||
4641 failed(verifyConvOutputSize(*this, IW, KW, outputShape.getDimSize(2),
4642 padValues[2], padValues[3], strideValues[1],
4643 dilationValues[1], "width", "x", "left",
4644 "right")))
4645 return failure();
4646 }
4647
4648 // Verify bias
4649 const ShapeAdaptor biasShape(getBias().getType());
4650 if (biasShape.hasRank() && outputShape.hasRank()) {
4651 const int64_t biasChannels = biasShape.getDimSize(0);
4652 const int64_t outputChannels =
4653 outputShape.getDimSize(outputShape.getRank() - 1);
4654 if (biasChannels == ShapedType::kDynamic ||
4655 outputChannels == ShapedType::kDynamic)
4656 // Skip following checks if biasChannels or outputChannels is dynamic dim
4657 return success();
4658
4659 if (biasChannels != outputChannels && biasChannels != 1)
4660 return emitOpError(
4661 "bias channels expected to be equal to output channels (")
4662 << outputChannels << ") or 1, got " << biasChannels;
4663 }
4664
4665 return success();
4666}
4667
4668LogicalResult Conv3DOp::inferReturnTypeComponents(
4669 MLIRContext *context, ::std::optional<Location> location,
4670 Conv3DOp::Adaptor adaptor,
4671 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4672 return inferConvReturnTypeComponents(adaptor, inferredReturnShapes);
4673}
4674
4675LogicalResult Conv3DOp::verify() {
4676 if (verifyConvOp(*this).failed() || verifyConvOpModes(*this).failed() ||
4677 verifyConvOpErrorIf(*this).failed())
4678 return failure();
4679 return success();
4680}
4681
4682LogicalResult AvgPool2dOp::inferReturnTypeComponents(
4683 MLIRContext *context, ::std::optional<Location> location,
4684 AvgPool2dOp::Adaptor adaptor,
4685 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4686 ShapeAdaptor inputShape(adaptor.getInput().getType());
4687 const Properties &prop = adaptor.getProperties();
4688 return poolingInferReturnTypes(inputShape, prop.kernel, prop.stride, prop.pad,
4689 inferredReturnShapes);
4690}
4691
4692LogicalResult AvgPool2dAdaptiveOp::inferReturnTypeComponents(
4693 MLIRContext *context, ::std::optional<Location> location,
4694 AvgPool2dAdaptiveOp::Adaptor adaptor,
4695 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4696 ShapeAdaptor inputShape(adaptor.getInput().getType());
4697
4698 llvm::SmallVector<int64_t> kernelValues;
4699 llvm::SmallVector<int64_t> strideValues;
4700 llvm::SmallVector<int64_t> padValues;
4701 if (tosa::getConstShapeValues(adaptor.getKernel().getDefiningOp(),
4702 kernelValues) &&
4703 tosa::getConstShapeValues(adaptor.getStride().getDefiningOp(),
4704 strideValues) &&
4705 tosa::getConstShapeValues(adaptor.getPad().getDefiningOp(), padValues)) {
4706 return poolingInferReturnTypes(inputShape, kernelValues, strideValues,
4707 padValues, inferredReturnShapes);
4708 }
4709
4710 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4711 if (inputShape.hasRank()) {
4712 // Keep N & C as pooling only changes H & W.
4713 outputShape[0] = inputShape.getDimSize(0);
4714 outputShape[3] = inputShape.getDimSize(3);
4715 }
4716
4717 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4718 return success();
4719}
4720
4721LogicalResult MaxPool2dOp::inferReturnTypeComponents(
4722 MLIRContext *context, ::std::optional<Location> location,
4723 MaxPool2dOp::Adaptor adaptor,
4724 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4725 ShapeAdaptor inputShape(adaptor.getInput().getType());
4726 const Properties &prop = adaptor.getProperties();
4727 return poolingInferReturnTypes(inputShape, prop.kernel, prop.stride, prop.pad,
4728 inferredReturnShapes);
4729}
4730
4731LogicalResult MaxPool2dAdaptiveOp::inferReturnTypeComponents(
4732 MLIRContext *context, ::std::optional<Location> location,
4733 MaxPool2dAdaptiveOp::Adaptor adaptor,
4734 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4735 ShapeAdaptor inputShape(adaptor.getInput().getType());
4736
4737 llvm::SmallVector<int64_t> kernelValues;
4738 llvm::SmallVector<int64_t> strideValues;
4739 llvm::SmallVector<int64_t> padValues;
4740 if (tosa::getConstShapeValues(adaptor.getKernel().getDefiningOp(),
4741 kernelValues) &&
4742 tosa::getConstShapeValues(adaptor.getStride().getDefiningOp(),
4743 strideValues) &&
4744 tosa::getConstShapeValues(adaptor.getPad().getDefiningOp(), padValues)) {
4745 return poolingInferReturnTypes(inputShape, kernelValues, strideValues,
4746 padValues, inferredReturnShapes);
4747 }
4748
4749 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4750 if (inputShape.hasRank()) {
4751 outputShape[0] = inputShape.getDimSize(0);
4752 outputShape[3] = inputShape.getDimSize(3);
4753 }
4754 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4755 return success();
4756}
4757
4758LogicalResult MaxPool2dOp::verify() {
4759 if (failed(verifySameElementTypes(*this, /* intype = */ getInput().getType(),
4760 /* outType = */ getOutput().getType())))
4761 return failure();
4762
4763 if (failed(verifyPoolingOp(*this)))
4764 return failure();
4765
4766 return success();
4767}
4768
4769LogicalResult MaxPool2dAdaptiveOp::verify() {
4770 if (failed(verifySameElementTypes(*this, /* intype = */ getInput().getType(),
4771 /* outType = */ getOutput().getType())))
4772 return failure();
4773
4774 AdaptivePoolingConstShapeValues values;
4776
4777 if (failed(verifyPoolingOpImpl(getOperation(), values.kernel, values.stride,
4778 values.pad, getInput(), getOutput())))
4779 return failure();
4780
4781 return success();
4782}
4783
4784LogicalResult DepthwiseConv2DOp::inferReturnTypeComponents(
4785 MLIRContext *context, ::std::optional<Location> location,
4786 DepthwiseConv2DOp::Adaptor adaptor,
4787 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4788 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4789
4790 int64_t inputWidth = ShapedType::kDynamic;
4791 int64_t inputHeight = ShapedType::kDynamic;
4792 int64_t inputChannels = ShapedType::kDynamic;
4793
4794 int64_t weightWidth = ShapedType::kDynamic;
4795 int64_t weightHeight = ShapedType::kDynamic;
4796 int64_t depthChannels = ShapedType::kDynamic;
4797
4798 // Input shape describes input width/height and batch.
4799 ShapeAdaptor inputShape(adaptor.getInput().getType());
4800 if (inputShape.hasRank()) {
4801 outputShape[0] = inputShape.getDimSize(0);
4802 inputHeight = inputShape.getDimSize(1);
4803 inputWidth = inputShape.getDimSize(2);
4804 inputChannels = inputShape.getDimSize(3);
4805 }
4806
4807 // Weight shapes describes the filter width/height and the output channels.
4808 ShapeAdaptor weightShape(adaptor.getWeight().getType());
4809 if (weightShape.hasRank()) {
4810 weightHeight = weightShape.getDimSize(0);
4811 weightWidth = weightShape.getDimSize(1);
4812 inputChannels = ShapedType::isDynamic(inputChannels)
4813 ? weightShape.getDimSize(2)
4814 : inputChannels;
4815 depthChannels = weightShape.getDimSize(3);
4816 }
4817
4818 // If both inputChannels and depthChannels are available we can determine
4819 // the output channels.
4820 if (ShapedType::isStatic(inputChannels) &&
4821 ShapedType::isStatic(depthChannels)) {
4822 outputShape[3] = inputChannels * depthChannels;
4823 }
4824
4825 // Bias shape can describe the output channels.
4826 ShapeAdaptor biasShape(adaptor.getBias().getType());
4827 if (biasShape.hasRank() && ShapedType::isDynamic(outputShape[3])) {
4828 int64_t bc = biasShape.getDimSize(0);
4829 if (bc != ShapedType::kDynamic && bc != 1)
4830 outputShape[3] = bc;
4831 }
4832
4833 llvm::ArrayRef<int64_t> dilation = adaptor.getDilation();
4834 llvm::ArrayRef<int64_t> padding = adaptor.getPad();
4835 llvm::ArrayRef<int64_t> stride = adaptor.getStride();
4836
4837 if (ShapedType::isStatic(inputHeight) && ShapedType::isStatic(weightHeight)) {
4838 int64_t inputSize = inputHeight + padding[0] + padding[1];
4839 int64_t filterSize = (weightHeight - 1) * dilation[0] + 1;
4840 int64_t unstridedResult = inputSize - filterSize + 1;
4841 outputShape[1] = (unstridedResult - 1) / stride[0] + 1;
4842 }
4843
4844 if (ShapedType::isStatic(inputWidth) && ShapedType::isStatic(weightWidth)) {
4845 int64_t inputSize = inputWidth + padding[2] + padding[3];
4846 int64_t filterSize = (weightWidth - 1) * dilation[1] + 1;
4847 int64_t unstridedResult = inputSize - filterSize + 1;
4848 outputShape[2] = (unstridedResult - 1) / stride[1] + 1;
4849 }
4850
4851 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4852 return success();
4853}
4854
4855LogicalResult DepthwiseConv2DOp::verify() {
4856 if (verifyConvOp(*this).failed() || verifyConvOpModes(*this).failed() ||
4857 verifyConvOpErrorIf(*this).failed())
4858 return failure();
4859 return success();
4860}
4861
4862LogicalResult TransposeConv2DOp::inferReturnTypeComponents(
4863 MLIRContext *context, ::std::optional<Location> location,
4864 TransposeConv2DOp::Adaptor adaptor,
4865 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4866 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4867
4868 int64_t inputWidth = ShapedType::kDynamic;
4869 int64_t inputHeight = ShapedType::kDynamic;
4870 int64_t weightWidth = ShapedType::kDynamic;
4871 int64_t weightHeight = ShapedType::kDynamic;
4872
4873 // Input shape describes input width/height and batch.
4874 ShapeAdaptor inputShape(adaptor.getInput().getType());
4875 if (inputShape.hasRank()) {
4876 outputShape[0] = ShapedType::isDynamic(outputShape[0])
4877 ? inputShape.getDimSize(0)
4878 : outputShape[0];
4879 inputHeight = inputShape.getDimSize(1);
4880 inputWidth = inputShape.getDimSize(2);
4881 }
4882
4883 // Weight shapes describes the filter width/height and the output channels.
4884 ShapeAdaptor weightShape(adaptor.getWeight().getType());
4885 if (weightShape.hasRank()) {
4886 outputShape[3] = ShapedType::isDynamic(outputShape[3])
4887 ? weightShape.getDimSize(0)
4888 : outputShape[3];
4889 weightHeight = weightShape.getDimSize(1);
4890 weightWidth = weightShape.getDimSize(2);
4891 }
4892
4893 // Bias shape can describe the output channels.
4894 ShapeAdaptor biasShape(adaptor.getBias().getType());
4895 if (biasShape.hasRank() && ShapedType::isDynamic(outputShape[3])) {
4896 int64_t bc = biasShape.getDimSize(0);
4897 if (bc != ShapedType::kDynamic && bc != 1)
4898 outputShape[3] = bc;
4899 }
4900
4901 llvm::ArrayRef<int64_t> padding = adaptor.getOutPad();
4902 llvm::ArrayRef<int64_t> stride = adaptor.getStride();
4903
4904 if (ShapedType::isStatic(inputHeight) && ShapedType::isStatic(weightHeight)) {
4905 int64_t calculateSize =
4906 (inputHeight - 1) * stride[0] + padding[0] + padding[1] + weightHeight;
4907 outputShape[1] =
4908 ShapedType::isDynamic(outputShape[1]) ? calculateSize : outputShape[1];
4909 }
4910
4911 if (ShapedType::isStatic(inputWidth) && ShapedType::isStatic(weightWidth)) {
4912 int64_t calculateSize =
4913 (inputWidth - 1) * stride[1] + padding[2] + padding[3] + weightWidth;
4914 outputShape[2] =
4915 ShapedType::isDynamic(outputShape[2]) ? calculateSize : outputShape[2];
4916 }
4917
4918 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4919 return success();
4920}
4921
4922LogicalResult TransposeConv2DOp::verify() {
4923 if (verifyConvOp(*this).failed() || verifyConvOpModes(*this).failed())
4924 return failure();
4925
4926 const llvm::ArrayRef<int64_t> strides = getStride();
4927 const int64_t strideY = strides[0];
4928 const int64_t strideX = strides[1];
4929
4930 if (strideY < 1 || strideX < 1)
4931 return emitOpError("expect all stride values to be >= 1, got [")
4932 << strides << "]";
4933
4934 const auto checkPadAgainstKernelDim =
4935 [this](int64_t padValue, int64_t kernelDimSize, llvm::StringRef padName,
4936 llvm::StringRef kernelDimName) -> LogicalResult {
4937 if (padValue <= -kernelDimSize)
4938 return emitOpError("expected ")
4939 << padName << " > -" << kernelDimName << ", but got: " << padName
4940 << "=" << padValue << " and " << kernelDimName << "="
4941 << kernelDimSize;
4942 return success();
4943 };
4944
4945 const llvm::ArrayRef<int64_t> padding = getOutPad();
4946 const int64_t outPadTop = padding[0];
4947 const int64_t outPadBottom = padding[1];
4948 const int64_t outPadLeft = padding[2];
4949 const int64_t outPadRight = padding[3];
4950
4951 const auto weightType =
4952 llvm::dyn_cast<RankedTensorType>(getWeight().getType());
4953
4954 if (weightType) {
4955 const int64_t kernelHeight = weightType.getDimSize(1);
4956 if (ShapedType::isStatic(kernelHeight)) {
4957 if (failed(checkPadAgainstKernelDim(outPadTop, kernelHeight,
4958 "out_pad_top", "KH")))
4959 return failure();
4960
4961 if (failed(checkPadAgainstKernelDim(outPadBottom, kernelHeight,
4962 "out_pad_bottom", "KH")))
4963 return failure();
4964 }
4965
4966 const int64_t kernelWidth = weightType.getDimSize(2);
4967 if (ShapedType::isStatic(kernelWidth)) {
4968 if (failed(checkPadAgainstKernelDim(outPadLeft, kernelWidth,
4969 "out_pad_left", "KW")))
4970 return failure();
4971
4972 if (failed(checkPadAgainstKernelDim(outPadRight, kernelWidth,
4973 "out_pad_right", "KW")))
4974 return failure();
4975 }
4976 }
4977
4978 // Rest of the checks depend on the output type being a RankedTensorType
4979 const auto outputType =
4980 llvm::dyn_cast<RankedTensorType>(getOutput().getType());
4981 if (!outputType)
4982 return success();
4983
4984 const auto inputType = llvm::dyn_cast<RankedTensorType>(getInput().getType());
4985 if (inputType && weightType) {
4986 const int64_t inputHeight = inputType.getDimSize(1);
4987 const int64_t kernelHeight = weightType.getDimSize(1);
4988 const int64_t outputHeight = outputType.getDimSize(1);
4989
4990 if (ShapedType::isStatic(inputHeight) &&
4991 ShapedType::isStatic(outputHeight)) {
4992 if (outputHeight !=
4993 (inputHeight - 1) * strideY + outPadTop + outPadBottom + kernelHeight)
4994 return emitOpError(
4995 "dimension mismatch: expected OH == (IH - 1) * stride_y "
4996 "+ out_pad_top + out_pad_bottom + KH, but got ")
4997 << outputHeight << " != (" << inputHeight << " - 1) * "
4998 << strideY << " + " << outPadTop << " + " << outPadBottom
4999 << " + " << kernelHeight;
5000 }
5001
5002 const int64_t inputWidth = inputType.getDimSize(2);
5003 const int64_t kernelWidth = weightType.getDimSize(2);
5004 const int64_t outputWidth = outputType.getDimSize(2);
5005
5006 if (ShapedType::isStatic(inputWidth) && ShapedType::isStatic(outputWidth)) {
5007 if (outputWidth !=
5008 (inputWidth - 1) * strideX + outPadLeft + outPadRight + kernelWidth)
5009 return emitOpError(
5010 "dimension mismatch: expected OW == (IW - 1) * stride_x "
5011 "+ out_pad_left + out_pad_right + KW, but got ")
5012 << outputWidth << " != (" << inputWidth << " - 1) * " << strideX
5013 << " + " << outPadLeft << " + " << outPadRight << " + "
5014 << kernelWidth;
5015 }
5016 }
5017
5018 const auto biasType = llvm::dyn_cast<RankedTensorType>(getBias().getType());
5019
5020 if (!biasType)
5021 return success();
5022
5023 const int64_t biasChannels = biasType.getDimSize(0);
5024
5025 // Skip further checks if bias is dynamic
5026 if (biasChannels == ShapedType::kDynamic)
5027 return success();
5028
5029 const int64_t outputChannels = outputType.getDimSize(3);
5030 if (!ShapedType::isDynamic(outputChannels) &&
5031 biasChannels != outputChannels && biasChannels != 1)
5032 return emitOpError(
5033 "bias channels expected to be equal to output channels (")
5034 << outputChannels << ") or 1, got " << biasChannels;
5035
5036 return success();
5037}
5038
5039LogicalResult RescaleOp::verify() {
5040 const auto inputType = llvm::cast<ShapedType>(getInput().getType());
5041 auto inputElementType =
5042 getStorageElementTypeOrSelf(inputType.getElementType());
5043 if (!mlir::isa<IntegerType>(inputElementType)) {
5044 emitOpError("expect input to have integer element type, got ")
5045 << inputElementType;
5046 return failure();
5047 }
5048
5049 const auto outputType = llvm::cast<ShapedType>(getOutput().getType());
5050 auto outputElementType =
5051 getStorageElementTypeOrSelf(outputType.getElementType());
5052 if (!mlir::isa<IntegerType>(outputElementType)) {
5053 emitOpError("expect output to have integer element type, got ")
5054 << outputElementType;
5055 return failure();
5056 }
5057
5058 if (verifyRescaleValueAndZpTypes(*this, getInput(), getInputZp(), "input")
5059 .failed())
5060 return failure();
5061
5062 if (verifyRescaleValueAndZpTypes(*this, getOutput(), getOutputZp(), "output")
5063 .failed())
5064 return failure();
5065
5066 FailureOr<int64_t> maybeIZp = getInputZeroPoint();
5067 if (succeeded(maybeIZp) && verifyInputZeroPoint(*maybeIZp).failed())
5068 return failure();
5069
5070 FailureOr<int64_t> maybeOZp = getOutputZeroPoint();
5071 if (succeeded(maybeOZp) && verifyOutputZeroPoint(*maybeOZp).failed())
5072 return failure();
5073
5074 const auto multiplierType = llvm::cast<ShapedType>(getMultiplier().getType());
5075 // multiplier element type must be i32 for scale32 = true
5076 if (getScale32() && !multiplierType.getElementType().isInteger(32)) {
5077 emitOpError("expect i32 element type for multiplier for scale32=true, got ")
5078 << multiplierType.getElementType();
5079 return failure();
5080 }
5081
5082 // multiplier element type must be i16 for scale32 = false
5083 if (!getScale32() && !multiplierType.getElementType().isInteger(16)) {
5084 emitOpError(
5085 "expect i16 element type for multiplier for scale32=false, got ")
5086 << multiplierType.getElementType();
5087 return failure();
5088 }
5089
5090 if (!inputType.hasRank())
5091 return success();
5092
5093 // multiplier/shift must have shape = {numChannels},
5094 // where numChannel is 1 if per_channel = false
5095 // otherwise numChannel is dimension in input shape's last axis
5096 int64_t numChannels = 1;
5097 if (getPerChannel()) {
5098 if (inputType.getRank() < 1) {
5099 emitOpError("requires input to be at least rank 1 when per_channel is "
5100 "true, but got rank ")
5101 << inputType.getRank();
5102 return failure();
5103 }
5104 numChannels = inputType.getDimSize(inputType.getRank() - 1);
5105 }
5106
5107 if (outputType.hasRank()) {
5109 getOperation(), outputType, inputType.getShape())))
5110 return failure();
5111 }
5112
5113 if (multiplierType.hasRank()) {
5114 ArrayRef<int64_t> multiplierShape = multiplierType.getShape();
5115 // multiplier input has rank 1 by dialect definition
5116 if (multiplierShape[0] != ShapedType::kDynamic &&
5117 multiplierShape[0] != numChannels) {
5118 emitOpError("expect shape of { ")
5119 << numChannels << " } for multiplier input, got { "
5120 << multiplierShape[0] << " }";
5121 return failure();
5122 }
5123 }
5124
5125 const auto shiftType = llvm::cast<ShapedType>(getShift().getType());
5126 if (shiftType.hasRank()) {
5127 ArrayRef<int64_t> shiftShape = shiftType.getShape();
5128 // shift input has rank 1 by dialect definition
5129 if (shiftShape[0] != ShapedType::kDynamic && shiftShape[0] != numChannels) {
5130 emitOpError("expect shape of { ")
5131 << numChannels << " } for shift input, got { " << shiftShape[0]
5132 << " }";
5133 return failure();
5134 }
5135 }
5136
5137 return success();
5138}
5139
5140LogicalResult RescaleOp::inferReturnTypeComponents(
5141 MLIRContext *context, ::std::optional<Location> location,
5142 RescaleOp::Adaptor adaptor,
5143 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5144 ShapeAdaptor inputShape(adaptor.getInput().getType());
5145 inferredReturnShapes.push_back(ShapedTypeComponents(inputShape));
5146 return success();
5147}
5148
5149LogicalResult CastOp::verify() {
5150 const ShapedType inputType = llvm::cast<ShapedType>(getInput().getType());
5151 const ShapedType outputType = llvm::cast<ShapedType>(getType());
5152 const Type inputElementType = inputType.getElementType();
5153 const Type outputElementType = outputType.getElementType();
5154
5155 const bool inputIsBlockScaled = llvm::isa<BlockScaledType>(inputElementType);
5156 const bool outputIsBlockScaled =
5157 llvm::isa<BlockScaledType>(outputElementType);
5158
5159 const bool isUnsigned = this->getInputUnsigned();
5160 const Type inputDataType = getStorageElementTypeOrSelf(inputType);
5161
5162 if (isUnsigned)
5163 if (!inputDataType.isInteger() || inputDataType.isInteger(1))
5164 return emitOpError()
5165 << "attribute input_unsigned requires integer type inputs. Got: "
5166 << inputDataType;
5167
5168 if (!inputIsBlockScaled && !outputIsBlockScaled)
5169 return success();
5170
5171 if (inputIsBlockScaled && outputIsBlockScaled)
5172 return emitOpError()
5173 << "requires exactly one of input or output to have block scaled "
5174 "element type";
5175
5176 const Type scalarElementType =
5177 inputIsBlockScaled ? outputElementType : inputElementType;
5178 if (!llvm::isa<FloatType>(scalarElementType))
5179 return emitOpError()
5180 << "requires non-block-scaled element type to be floating-point "
5181 "when casting to or from block scaled element type, got "
5182 << scalarElementType;
5183
5184 return success();
5185}
5186
5187LogicalResult CastFromBlockScaledOp::inferReturnTypeComponents(
5188 MLIRContext *context, ::std::optional<Location> location,
5189 CastFromBlockScaledOp::Adaptor adaptor,
5190 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5191 const ShapeAdaptor inputShape(adaptor.getInputData().getType());
5192 inferredReturnShapes.push_back(ShapedTypeComponents(inputShape));
5193 return success();
5194}
5195
5196LogicalResult CastFromBlockScaledOp::verify() {
5197 const Type inputDataType = getInputData().getType();
5198 const Type outputDataType = getResult().getType();
5199 if (failed(verifyCompatibleShape(inputDataType, outputDataType)))
5200 return emitOpError() << "require compatible shapes for input_data ("
5201 << inputDataType << ") and " << "output_data ("
5202 << outputDataType << ")";
5203
5204 const ShapeAdaptor inputDataShape = ShapeAdaptor(inputDataType);
5205
5206 if (inputDataShape.hasRank()) {
5207 const unsigned int blockSize =
5208 BlockSizeAttr::getBlockSizeValue(getBlockSize());
5209 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
5210 return emitOpError("expect block size to be 32, got ") << blockSize;
5211 const int64_t inputDataLastDim =
5212 inputDataShape.getDimSize(inputDataShape.getRank() - 1);
5213 if (inputDataLastDim % blockSize != 0)
5214 return emitOpError() << "expect last dimension of input_data ("
5215 << inputDataLastDim
5216 << ") to be divisible by block_size (" << blockSize
5217 << ")";
5218
5219 const Type inputScaleType = getInputScale().getType();
5220 const ShapeAdaptor inputScaleShape = ShapeAdaptor(inputScaleType);
5221
5222 if (inputScaleShape.hasRank()) {
5223 SmallVector<int64_t> inputDataDims, inputScaleDims;
5224 inputDataShape.getDims(inputDataDims);
5225 inputScaleShape.getDims(inputScaleDims);
5226
5227 if (inputDataDims.size() != inputScaleDims.size() ||
5229 ArrayRef<int64_t>(inputDataDims).drop_back(1),
5230 ArrayRef<int64_t>(inputScaleDims).drop_back(1))))
5231 return emitOpError()
5232 << "require compatible shapes for input_data (" << inputDataType
5233 << ") and " << "input_scale (" << inputScaleType
5234 << ") except for the last dimension";
5235
5236 const SmallVector<int64_t, 2> dimsToCheck{inputDataLastDim / blockSize,
5237 inputScaleDims.back()};
5238 if (ShapedType::isStatic(inputDataLastDim) &&
5239 failed(verifyCompatibleDims(dimsToCheck)))
5240 return emitOpError()
5241 << "expect last dimension of input_scale ("
5242 << inputScaleDims.back()
5243 << ") to be equal to last dimension of input_data / block_size ("
5244 << inputDataDims.back() / blockSize << ")";
5245 }
5246 }
5247
5248 return success();
5249}
5250
5251LogicalResult CastToBlockScaledOp::inferReturnTypeComponents(
5252 MLIRContext *context, ::std::optional<Location> location,
5253 CastToBlockScaledOp::Adaptor adaptor,
5254 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5255 const ShapeAdaptor inputShape(adaptor.getInputData().getType());
5256 inferredReturnShapes.push_back(ShapedTypeComponents(inputShape));
5257 if (!inputShape.hasRank())
5258 return success();
5259
5260 // Calculate output_scale shape if ranked input provided
5261 SmallVector<int64_t> outputScaleShape;
5262 inputShape.getDims(outputScaleShape);
5263 const int64_t lastDimLoc = inputShape.getRank() - 1;
5264 const int64_t lastDimSize = inputShape.getDimSize(lastDimLoc);
5265 if (ShapedType::isStatic(lastDimSize)) {
5266 const unsigned int blockSize =
5267 BlockSizeAttr::getBlockSizeValue(adaptor.getBlockSize());
5268 outputScaleShape[lastDimLoc] = lastDimSize / blockSize;
5269 }
5270 inferredReturnShapes.push_back(ShapedTypeComponents(outputScaleShape));
5271 return success();
5272}
5273
5274LogicalResult CastToBlockScaledOp::verify() {
5275 const Type inputDataType = getInputData().getType();
5276 const Type outputDataType = getResult(0).getType();
5277 if (failed(verifyCompatibleShape(inputDataType, outputDataType)))
5278 return emitOpError() << "require compatible shapes for input_data ("
5279 << inputDataType << ") and " << "output_data ("
5280 << outputDataType << ")";
5281
5282 const unsigned int blockSize =
5283 BlockSizeAttr::getBlockSizeValue(getBlockSize());
5284 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
5285 return emitOpError("expect block size to be 32, got ") << blockSize;
5286 const ShapeAdaptor inputDataShape = ShapeAdaptor(inputDataType);
5287 if (inputDataShape.hasRank()) {
5288 const int64_t inputDataLastDim =
5289 inputDataShape.getDimSize(inputDataShape.getRank() - 1);
5290 if (ShapedType::isStatic(inputDataLastDim) &&
5291 inputDataLastDim % blockSize != 0)
5292 return emitOpError() << "expect last dimension of input_data ("
5293 << inputDataLastDim
5294 << ") to be divisible by block_size (" << blockSize
5295 << ")";
5296 }
5297
5298 const ShapeAdaptor outputDataShape = ShapeAdaptor(outputDataType);
5299 const Type outputScaleType = getResult(1).getType();
5300 const ShapeAdaptor outputScaleShape = ShapeAdaptor(outputScaleType);
5301 if (outputDataShape.hasRank() && outputScaleShape.hasRank()) {
5302 SmallVector<int64_t> outputDataDims, outputScaleDims;
5303 outputDataShape.getDims(outputDataDims);
5304 outputScaleShape.getDims(outputScaleDims);
5305
5306 if (outputDataDims.size() != outputScaleDims.size() ||
5308 ArrayRef<int64_t>(outputDataDims).drop_back(1),
5309 ArrayRef<int64_t>(outputScaleDims).drop_back(1))))
5310 return emitOpError() << "require compatible shapes for output_data ("
5311 << outputDataType << ") and " << "output_scale ("
5312 << outputScaleType
5313 << ") except for the last dimension";
5314
5315 const int64_t outputDataLastDim = outputDataDims.back();
5316 const SmallVector<int64_t, 2> dimsToCheck{outputDataLastDim / blockSize,
5317 outputScaleDims.back()};
5318 if (ShapedType::isStatic(outputDataLastDim) &&
5319 failed(verifyCompatibleDims(dimsToCheck)))
5320 return emitOpError()
5321 << "expect last dimension of output_scale ("
5322 << outputScaleDims.back()
5323 << ") to be equal to last dimension of output_data / block_size ("
5324 << outputDataDims.back() / blockSize << ")";
5325 }
5326
5327 return success();
5328}
5329
5330LogicalResult IfOp::inferReturnTypeComponents(
5331 MLIRContext *context, ::std::optional<Location> location,
5332 IfOp::Adaptor adaptor,
5333 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5334 llvm::SmallVector<tosa::YieldOp> yieldOps;
5335 for (Region *region : adaptor.getRegions()) {
5336 for (auto &block : *region)
5337 if (auto returnOp = dyn_cast<tosa::YieldOp>(block.getTerminator()))
5338 yieldOps.push_back(returnOp);
5339 }
5340
5341 if (yieldOps.empty())
5342 return failure();
5343
5344 // Get the initial type information for the yield op.
5345 llvm::SmallVector<ValueKnowledge> resultKnowledge;
5346 resultKnowledge.reserve(yieldOps.front().getNumOperands());
5347 for (auto operand : yieldOps.front().getOperands()) {
5348 resultKnowledge.push_back(
5349 ValueKnowledge::getKnowledgeFromType(operand.getType()));
5350 }
5351
5352 for (auto yieldOp : yieldOps) {
5353 if (resultKnowledge.size() != yieldOp.getNumOperands())
5354 return failure();
5355
5356 for (const auto &it : llvm::enumerate(yieldOp.getOperands())) {
5357 int32_t index = it.index();
5358 auto meet = ValueKnowledge::meet(
5359 resultKnowledge[index],
5360 ValueKnowledge::getKnowledgeFromType(it.value().getType()));
5361 if (!meet)
5362 continue;
5363 resultKnowledge[index] = meet;
5364 }
5365 }
5366
5367 for (const ValueKnowledge &result : resultKnowledge) {
5368 inferredReturnShapes.push_back(result.getShapedTypeComponents());
5369 }
5370
5371 return success();
5372}
5373
5374LogicalResult WhileOp::inferReturnTypeComponents(
5375 MLIRContext *context, ::std::optional<Location> location,
5376 WhileOp::Adaptor adaptor,
5377 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5378 llvm::SmallVector<tosa::YieldOp> yieldOps;
5379 for (auto &block : adaptor.getBodyGraph())
5380 if (auto returnOp = dyn_cast<tosa::YieldOp>(block.getTerminator()))
5381 yieldOps.push_back(returnOp);
5382
5383 // TOSA's while must have a tosa.yield as its terminator. If not found this
5384 // tosa.while is invalid.
5385 if (yieldOps.empty())
5386 return failure();
5387
5388 // Get the initial type information from the operand types.
5389 llvm::SmallVector<ValueKnowledge> resultKnowledge;
5390 resultKnowledge.reserve(yieldOps.front().getNumOperands());
5391 for (auto operand : yieldOps.front().getOperands()) {
5392 resultKnowledge.push_back(
5393 ValueKnowledge::getKnowledgeFromType(operand.getType()));
5394 }
5395
5396 for (auto yieldOp : yieldOps) {
5397 if (resultKnowledge.size() != yieldOp.getNumOperands())
5398 return failure();
5399
5400 for (const auto &it : llvm::enumerate(yieldOp.getOperands())) {
5401 int32_t index = it.index();
5402 if (auto meet = ValueKnowledge::meet(
5403 resultKnowledge[index],
5404 ValueKnowledge::getKnowledgeFromType(it.value().getType()))) {
5405 resultKnowledge[index] = meet;
5406 }
5407 }
5408 }
5409
5410 for (const ValueKnowledge &result : resultKnowledge) {
5411 inferredReturnShapes.push_back(result.getShapedTypeComponents());
5412 }
5413
5414 return success();
5415}
5416
5417std::optional<SmallVector<int64_t, 4>> ApplyScaleOp::getShapeForUnroll() {
5418 if (auto vt = llvm::dyn_cast<VectorType>(getType()))
5419 return llvm::to_vector<4>(vt.getShape());
5420 return std::nullopt;
5421}
5422
5424 Block::BlockArgListType blocksArgs,
5425 ValueRange initializers,
5426 StringRef prefix = "") {
5427 assert(blocksArgs.size() == initializers.size() &&
5428 "expected same length of arguments and initializers");
5429 if (initializers.empty())
5430 return;
5431
5432 parser << prefix << '(';
5433 llvm::interleaveComma(
5434 llvm::zip(blocksArgs, initializers), parser,
5435 [&](auto it) { parser << std::get<0>(it) << " = " << std::get<1>(it); });
5436 parser << ")";
5437}
5438
5439// parse and print of IfOp refer to the implementation of SCF dialect.
5440ParseResult IfOp::parse(OpAsmParser &parser, OperationState &result) {
5441 // Create the regions for 'then'.
5442 result.regions.reserve(2);
5443 Region *thenRegion = result.addRegion();
5444 Region *elseRegion = result.addRegion();
5445
5446 OpAsmParser::UnresolvedOperand cond;
5447
5448 if (parser.parseOperand(cond))
5449 return failure();
5450
5451 SmallVector<OpAsmParser::Argument, 4> regionArgs;
5452 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
5453
5454 // Parse the optional block arguments
5455 OptionalParseResult listResult =
5456 parser.parseOptionalAssignmentList(regionArgs, operands);
5457 if (listResult.has_value() && failed(listResult.value()))
5458 return failure();
5459
5460 // Parse a colon.
5461 if (failed(parser.parseColon()))
5462 return parser.emitError(parser.getCurrentLocation(),
5463 "expected type for condition operand");
5464
5465 // Parse the type of the condition operand
5466 Type condType;
5467 if (failed(parser.parseType(condType)))
5468 return parser.emitError(parser.getCurrentLocation(),
5469 "expected type for condition operand");
5470
5471 // Resolve operand with provided type
5472 if (failed(parser.resolveOperand(cond, condType, result.operands)))
5473 return failure();
5474
5475 // Parse optional block arg types
5476 if (listResult.has_value()) {
5477 FunctionType functionType;
5478
5479 if (failed(parser.parseType(functionType)))
5480 return parser.emitError(parser.getCurrentLocation())
5481 << "expected list of types for block arguments "
5482 << "followed by arrow type and list of return types";
5483
5484 result.addTypes(functionType.getResults());
5485
5486 if (functionType.getNumInputs() != operands.size()) {
5487 return parser.emitError(parser.getCurrentLocation())
5488 << "expected as many input types as operands " << "(expected "
5489 << operands.size() << " got " << functionType.getNumInputs()
5490 << ")";
5491 }
5492
5493 // Resolve input operands.
5494 if (failed(parser.resolveOperands(operands, functionType.getInputs(),
5495 parser.getCurrentLocation(),
5496 result.operands)))
5497 return failure();
5498 } else {
5499 // Parse optional results type list.
5500 if (parser.parseOptionalArrowTypeList(result.types))
5501 return failure();
5502 }
5503
5504 // Parse the 'then' region.
5505 if (parser.parseRegion(*thenRegion, /*arguments=*/{}, /*argTypes=*/{}))
5506 return failure();
5507
5508 // If we find an 'else' keyword then parse the 'else' region.
5509 if (!parser.parseOptionalKeyword("else")) {
5510 if (parser.parseRegion(*elseRegion, /*arguments=*/{}, /*argTypes=*/{}))
5511 return failure();
5512 }
5513
5514 // Parse the optional attribute list.
5515 if (parser.parseOptionalAttrDict(result.attributes))
5516 return failure();
5517 return success();
5518}
5519
5520void IfOp::print(OpAsmPrinter &p) {
5521 p << " " << getCondition();
5522
5523 printInitializationList(p, getThenGraph().front().getArguments(),
5524 getInputList(), " ");
5525 p << " : ";
5526 p << getCondition().getType();
5527
5528 if (!getInputList().empty()) {
5529 p << " (";
5530 llvm::interleaveComma(getInputList().getTypes(), p);
5531 p << ")";
5532 }
5533 p.printArrowTypeList(getResultTypes());
5534 p << " ";
5535
5536 p.printRegion(getThenGraph());
5537
5538 // Print the 'else' regions if it exists and has a block.
5539 auto &elseRegion = getElseGraph();
5540 if (!elseRegion.empty()) {
5541 p << " else ";
5542 p.printRegion(elseRegion);
5543 }
5544
5545 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
5546}
5547
5548LogicalResult IfOp::verify() {
5549 if (errorIfTypeOrShapeMismatch(*this, getThenGraph().front().getArguments(),
5550 "'then_graph' arguments", getInputList(),
5551 "'input_list'")
5552 .failed())
5553 return failure();
5554
5555 if (errorIfTypeOrShapeMismatch(*this, getElseGraph().front().getArguments(),
5556 "'else_graph' arguments", getInputList(),
5557 "'input_list'")
5558 .failed())
5559 return failure();
5560
5561 // MLIR will verify the absence of the terminator for us if otherwise.
5562 if (getThenGraph().front().mightHaveTerminator()) {
5563 auto thenYield =
5564 dyn_cast<tosa::YieldOp>(getThenGraph().front().getTerminator());
5565 if (thenYield && errorIfTypeOrShapeMismatch(
5566 *this, thenYield.getInputs(), "'then_graph' results",
5567 getOutputList(), "'output_list'")
5568 .failed())
5569 return failure();
5570 }
5571
5572 // MLIR will verify the absence of the terminator for us if otherwise.
5573 if (getElseGraph().front().mightHaveTerminator()) {
5574 auto elseYield =
5575 dyn_cast<tosa::YieldOp>(getElseGraph().front().getTerminator());
5576 if (elseYield && errorIfTypeOrShapeMismatch(
5577 *this, elseYield.getInputs(), "'else_graph' results",
5578 getOutputList(), "'output_list'")
5579 .failed())
5580 return failure();
5581 }
5582
5583 auto condType = getCondition().getType();
5584 if (errorIfShapeNotSizeOne(*this, condType).failed())
5585 return emitOpError() << "'condition' must be a size 1 tensor, got "
5586 << condType;
5587
5588 return success();
5589}
5590
5591LogicalResult WhileOp::verify() {
5592 if (errorIfTypeOrShapeMismatch(*this, getInputList(), "'input_list'",
5593 getOutputList(), "'output_list'")
5594 .failed())
5595 return failure();
5596
5597 if (errorIfTypeOrShapeMismatch(*this, getCondGraph().front().getArguments(),
5598 "'cond_graph' arguments", getInputList(),
5599 "'input_list'")
5600 .failed())
5601 return failure();
5602
5603 if (errorIfTypeOrShapeMismatch(*this, getBodyGraph().front().getArguments(),
5604 "'body_graph' arguments", getInputList(),
5605 "'input_list'")
5606 .failed())
5607 return failure();
5608
5609 if (getBodyGraph().front().mightHaveTerminator()) {
5610 auto bodyYield =
5611 dyn_cast<tosa::YieldOp>(getBodyGraph().front().getTerminator());
5612 if (bodyYield && errorIfTypeOrShapeMismatch(*this, bodyYield.getInputs(),
5613 "'body_graph' results",
5614 getInputList(), "'input_list'")
5615 .failed())
5616 return failure();
5617 }
5618
5619 // Condition block output must be a single element tensor with a single bool
5620 // value.
5621 if (!getCondGraph().front().mightHaveTerminator())
5622 return success();
5623
5624 auto condYield =
5625 dyn_cast<tosa::YieldOp>(getCondGraph().front().getTerminator());
5626 if (!condYield)
5627 return success();
5628
5629 if (condYield.getInputs().size() != 1)
5630 return emitOpError() << "require 'cond_graph' only have one result";
5631
5632 auto condOutType = condYield.getInputs()[0].getType();
5633 if (errorIfShapeNotSizeOne(*this, condOutType).failed())
5634 return emitOpError() << "'cond_graph' result must be a size 1 tensor, got "
5635 << condOutType;
5636
5637 if (!getElementTypeOrSelf(condOutType).isInteger(1))
5638 return emitOpError() << "'cond_graph' result must be a boolean tensor, got "
5639 << condOutType;
5640
5641 return success();
5642}
5643
5644LogicalResult ReverseOp::verify() {
5645 TensorType inputType = getInput1().getType();
5646 int32_t reverseAxis = getAxis();
5647
5648 if (reverseAxis < 0)
5649 return emitOpError("expected non-negative reverse axis");
5650 if (inputType.hasRank()) {
5651 int64_t inputRank = inputType.getRank();
5652 // We allow for a special case where the input/output shape has rank 0 and
5653 // axis is also 0.
5654 if (reverseAxis >= inputRank && (reverseAxis != 0 || inputRank != 0))
5655 return emitOpError("expect input tensor rank (")
5656 << inputRank << ") to be larger than reverse axis (" << reverseAxis
5657 << ")";
5658 }
5659
5660 return success();
5661}
5662
5663LogicalResult tosa::SelectOp::verify() {
5664 // verify input2 and input3 have same element type as output
5665 if (verifySameElementTypes(*this, /* inType = */ getOnTrue().getType(),
5666 /* outType = */ getOutput().getType())
5667 .failed() ||
5668 verifySameElementTypes(*this, /* inType = */ getOnFalse().getType(),
5669 /* outType = */ getOutput().getType())
5670 .failed()) {
5671 return failure();
5672 }
5673 // verify input1 has element type of bool
5674 auto predicateType = llvm::dyn_cast<ShapedType>(getPred().getType());
5675 if (!predicateType) {
5676 return emitOpError("expect shaped tensor for input1, got ")
5677 << getInput1().getType();
5678 }
5679 auto predicateElementType = predicateType.getElementType();
5680 if (!predicateElementType.isInteger(1)) {
5681 return emitOpError("expect element type of bool for input1, got ")
5682 << predicateElementType;
5683 }
5684
5685 return success();
5686}
5687
5688LogicalResult tosa::VariableReadOp::verify() {
5689 if (verifyVariableOpErrorIf(*this, getOutput1().getType(), "'output1'")
5690 .failed())
5691 return failure();
5692
5693 return success();
5694}
5695
5696LogicalResult tosa::VariableWriteOp::verify() {
5697 if (verifyVariableOpErrorIf(*this, getInput1().getType(), "'input1'")
5698 .failed())
5699 return failure();
5700
5701 return success();
5702}
5703
5704// parse and print of WhileOp refer to the implementation of SCF dialect.
5705ParseResult WhileOp::parse(OpAsmParser &parser, OperationState &result) {
5706 SmallVector<OpAsmParser::Argument, 4> regionArgs;
5707 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
5708 Region *cond = result.addRegion();
5709 Region *body = result.addRegion();
5710
5711 OptionalParseResult listResult =
5712 parser.parseOptionalAssignmentList(regionArgs, operands);
5713 if (listResult.has_value() && failed(listResult.value()))
5714 return failure();
5715
5716 FunctionType functionType;
5717 SMLoc typeLoc = parser.getCurrentLocation();
5718 if (failed(parser.parseColonType(functionType)))
5719 return failure();
5720
5721 result.addTypes(functionType.getResults());
5722
5723 if (functionType.getNumInputs() != operands.size()) {
5724 return parser.emitError(typeLoc)
5725 << "expected as many input types as operands " << "(expected "
5726 << operands.size() << " got " << functionType.getNumInputs() << ")";
5727 }
5728
5729 // Resolve input operands.
5730 if (failed(parser.resolveOperands(operands, functionType.getInputs(),
5731 parser.getCurrentLocation(),
5732 result.operands)))
5733 return failure();
5734
5735 // Propagate the types into the region arguments.
5736 for (size_t i = 0, e = regionArgs.size(); i != e; ++i)
5737 regionArgs[i].type = functionType.getInput(i);
5738
5739 return failure(parser.parseRegion(*cond, regionArgs) ||
5740 parser.parseKeyword("do") || parser.parseRegion(*body) ||
5741 parser.parseOptionalAttrDictWithKeyword(result.attributes));
5742}
5743
5744void WhileOp::print(OpAsmPrinter &parser) {
5745 printInitializationList(parser, getCondGraph().front().getArguments(),
5746 getInputList(), " ");
5747 parser << " : ";
5748 parser.printFunctionalType(getInputList().getTypes(),
5749 getResults().getTypes());
5750 parser << ' ';
5751 parser.printRegion(getCondGraph(), /*printEntryBlockArgs=*/false);
5752 parser << " do ";
5753 parser.printRegion(getBodyGraph());
5755 (*this)->getDiscardableAttrDictionary().getValue());
5756}
5757
5758// Create a rank-1 const tensor for zero point of the source tensor.
5759std::optional<Value> mlir::tosa::createZeroPointTensor(OpBuilder &builder,
5760 Location loc,
5761 Type srcElemType,
5762 int64_t zp) {
5763 srcElemType = getStorageElementTypeOrSelf(srcElemType);
5764 auto zpType = mlir::RankedTensorType::get({1}, srcElemType);
5765 if (llvm::isa<FloatType>(srcElemType)) {
5766 auto zpAttr = DenseElementsAttr::get(
5767 zpType, builder.getFloatAttr(srcElemType, static_cast<double>(zp)));
5768 return tosa::ConstOp::create(builder, loc, zpType, zpAttr);
5769 }
5770 if (llvm::isa<IntegerType>(srcElemType)) {
5771 auto zpAttr =
5772 DenseElementsAttr::get(zpType, builder.getIntegerAttr(srcElemType, zp));
5773 return tosa::ConstOp::create(builder, loc, zpType, zpAttr);
5774 }
5775 llvm::errs() << "zero point is not allowed for unsupported data types\n";
5776 return std::nullopt;
5777}
5778
5779//===----------------------------------------------------------------------===//
5780// TOSA Shape and Shape Operators Helper functions.
5781//===----------------------------------------------------------------------===//
5782
5784 return mlir::isa<tosa::shapeType>(t);
5785}
5786
5787LogicalResult
5788mlir::tosa::shapeType::verify(function_ref<InFlightDiagnostic()> emitError,
5789 int rank) {
5790 if (rank < 0)
5791 return emitError() << "invalid rank (must be >= 0): " << rank;
5792 return success();
5793}
5794
5796 for (auto v : op->getOperands()) {
5797 if (mlir::isa<::mlir::tosa::shapeType>(v.getType())) {
5798 Operation *definingOp = v.getDefiningOp();
5799 if (!definingOp || !definingOp->hasTrait<TosaShapeOperator>()) {
5800 return op->emitOpError("shape operand is not compile time resolvable");
5801 }
5802 }
5803 }
5804 return success();
5805}
5806
5807LogicalResult
5809 if (failed(OpTrait::impl::verifyAtLeastNOperands(op, 1)))
5810 return failure();
5811
5812 // delegate function that returns rank of shape type
5813 auto getRank = [](const Type type) {
5814 return mlir::cast<mlir::tosa::shapeType>(type).getRank();
5815 };
5816 auto operandTypes = op->getOperandTypes();
5817 auto resultTypes = op->getResultTypes();
5818
5819 auto rank = getRank(*op->getOperandTypes().begin());
5820 for (auto type : operandTypes) {
5821 if (getRank(type) != rank) {
5822 return op->emitOpError("operands don't have matching ranks");
5823 }
5824 }
5825 for (auto type : resultTypes) {
5826 if (getRank(type) != rank) {
5827 return op->emitOpError("result shape has different rank than operands");
5828 }
5829 }
5830 return success();
5831}
5832
5833//===----------------------------------------------------------------------===//
5834// TOSA Shape Operators verify functions.
5835//===----------------------------------------------------------------------===//
5836
5837LogicalResult tosa::ConstShapeOp::verify() {
5838 // check one dimensional rank
5839 auto valuesRank = getValues().getType().getRank();
5840 if (valuesRank != 1)
5841 return emitOpError("expect elements in attribute values with rank 1");
5842 // check that number of elements in values attr equal to rank of result shape
5843 auto count = getValues().getNumElements();
5844 auto rank = (cast<tosa::shapeType>(getResult().getType())).getRank();
5845 if (count != rank && (count != 1 || rank != 0)) {
5846 return emitOpError("expect number of elements in attribute values (")
5847 << count << ") to be equal to the rank (" << rank
5848 << ") for the result shape type";
5849 }
5850 return success();
5851}
5852
5853LogicalResult tosa::DimOp::verify() {
5854 const tosa::shapeType outShapeType =
5855 cast<tosa::shapeType>(getResult().getType());
5856 if (outShapeType.getRank() != 1)
5857 return emitOpError("expect output shape type to contain one element, got ")
5858 << outShapeType;
5859
5860 const ShapeAdaptor inputType(getInput1().getType());
5861 if (inputType.hasRank()) {
5862 const int64_t inputRank = inputType.getRank();
5863 const int64_t axis = getAxisAttr().getInt();
5864 if (axis < 0 || axis >= inputRank)
5865 return emitOpError("expect axis to be in the range [0, ")
5866 << inputRank << "), got " << axis;
5867 }
5868 return success();
5869}
5870
5871LogicalResult tosa::ConcatShapeOp::verify() {
5872 const tosa::shapeType outShapeType =
5873 cast<tosa::shapeType>(getResult().getType());
5874 const int64_t outputRank = outShapeType.getRank();
5875 const Operation::operand_range inputList = getInput();
5876
5877 if (inputList.size() == 0)
5878 return emitOpError("requires at least one input shape");
5879
5880 if (llvm::any_of(inputList, [](Value v) {
5881 return cast<tosa::shapeType>(v.getType()).getRank() == 0;
5882 }))
5883 return emitOpError("requires all inputs shapes have a rank greater than 0");
5884
5885 const int64_t inputsRank =
5886 llvm::accumulate(inputList, 0, [](int64_t acc, const Value &input) {
5887 const tosa::shapeType inShapeType =
5888 cast<tosa::shapeType>(input.getType());
5889 return acc + inShapeType.getRank();
5890 });
5891 if (outputRank != inputsRank)
5892 return emitOpError("requires output shape rank to be equal to the sum of "
5893 "the input shape ranks (")
5894 << inputsRank << "), got " << outputRank;
5895
5896 return success();
5897}
5898
5899LogicalResult tosa::SliceShapeOp::verify() {
5900 std::optional<int32_t> start;
5901 DenseIntElementsAttr startAttr;
5902 if (matchPattern(getStart(), m_Constant(&startAttr)))
5903 start = startAttr.getValues<int32_t>()[0];
5904 if (start && start.value() < 0)
5905 return emitOpError("expected non-negative start index, got ")
5906 << start.value();
5907
5908 std::optional<int32_t> size;
5909 DenseIntElementsAttr sizeAttr;
5910 if (matchPattern(getSize(), m_Constant(&sizeAttr)))
5911 size = sizeAttr.getValues<int32_t>()[0];
5912 if (size && size.value() <= 0)
5913 return emitOpError("expected positive size, got ") << size.value();
5914
5915 if (!size)
5916 return success();
5917
5918 const tosa::shapeType outShapeType =
5919 cast<tosa::shapeType>(getResult().getType());
5920 const int64_t outputRank = outShapeType.getRank();
5921 if (outputRank != size)
5922 return emitOpError(
5923 "expected output type size to be equal to size attribute, got ")
5924 << outputRank << " vs " << size.value();
5925
5926 if (!start)
5927 return success();
5928
5929 const tosa::shapeType inShapeType =
5930 cast<tosa::shapeType>(getInput().getType());
5931 const int64_t inputRank = inShapeType.getRank();
5932 const int64_t sliceSize = start.value() + size.value();
5933 if (sliceSize > inputRank)
5934 return emitOpError("expected start + size to be less than or equal to "
5935 "input shape rank (")
5936 << inputRank << "), got " << sliceSize;
5937
5938 return success();
5939}
5940
5941//===----------------------------------------------------------------------===//
5942// TOSA Attribute Definitions.
5943//===----------------------------------------------------------------------===//
5944
5945#define GET_ATTRDEF_CLASSES
5946#include "mlir/Dialect/Tosa/IR/TosaAttributes.cpp.inc"
5947
5948//===----------------------------------------------------------------------===//
5949// TOSA Type Definitions.
5950//===----------------------------------------------------------------------===//
5951#define GET_TYPEDEF_CLASSES
5952#include "mlir/Dialect/Tosa/IR/TosaOpsTypesBase.cpp.inc"
5953
5954//===----------------------------------------------------------------------===//
5955// TOSA Operator Definitions.
5956//===----------------------------------------------------------------------===//
5957
5958static ParseResult parseOptionalBoolClause(OpAsmParser &parser,
5959 StringRef keyword,
5960 BoolAttr &result) {
5961 if (failed(parser.parseOptionalKeyword(keyword)))
5962 return success();
5963 if (parser.parseLParen() || parser.parseAttribute(result) ||
5964 parser.parseRParen())
5965 return failure();
5966 return success();
5967}
5968
5969static void printOptionalBoolClause(OpAsmPrinter &printer, StringRef keyword,
5970 BoolAttr attr) {
5971 if (!attr)
5972 return;
5973 printer << keyword << '(';
5974 printer.printAttribute(attr);
5975 printer << ')';
5976}
5977
5978static ParseResult parseLocalBound(OpAsmParser &parser, BoolAttr &result) {
5979 return parseOptionalBoolClause(parser, "local_bound", result);
5980}
5981
5982static void printLocalBound(OpAsmPrinter &printer, Operation *, BoolAttr attr) {
5983 printOptionalBoolClause(printer, "local_bound", attr);
5984}
5985
5986static ParseResult parseInputUnsigned(OpAsmParser &parser, BoolAttr &result) {
5987 return parseOptionalBoolClause(parser, "input_unsigned", result);
5988}
5989
5991 BoolAttr attr) {
5992 printOptionalBoolClause(printer, "input_unsigned", attr);
5993}
5994
5995#define GET_OP_CLASSES
5996#include "mlir/Dialect/Tosa/IR/TosaOps.cpp.inc"
return success()
static void printInitializationList(OpAsmPrinter &p, Block::BlockArgListType blocksArgs, ValueRange initializers, StringRef prefix="")
Prints the initialization list in the form of <prefix>(inner = outer, inner2 = outer2,...
Definition SCF.cpp:502
true
Given two iterators into the same block, return "true" if a is before `b.
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
static Type getValueType(Attribute attr)
Definition SPIRVOps.cpp:835
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
Definition SPIRVOps.cpp:229
static LogicalResult verifyMatMulShapes(T op, bool transposeB)
Definition TosaOps.cpp:2072
static ParseResult parseOptionalBoolClause(OpAsmParser &parser, StringRef keyword, BoolAttr &result)
Definition TosaOps.cpp:5958
static void printShapeToDiagnostic(InFlightDiagnostic &diag, ArrayRef< int64_t > shape)
Definition TosaOps.cpp:356
static void buildMatMulOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value a, Value b)
Definition TosaOps.cpp:1375
static LogicalResult verifySameElementTypes(Operation *op, Type aType, Type bType, StringRef aName="input", StringRef bName="output")
Definition TosaOps.cpp:1003
static ParseResult parseLocalBound(OpAsmParser &parser, BoolAttr &result)
Definition TosaOps.cpp:5978
LogicalResult inferConvReturnTypeComponents(AdaptorT adaptor, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
Definition TosaOps.cpp:4457
static int64_t getMatMulBatchDim(const ShapeAdaptor &shape, int64_t outputRank, int64_t axis)
Definition TosaOps.cpp:1931
static SmallVector< int64_t > convertToMlirShape(ArrayRef< int64_t > shape)
Definition TosaOps.cpp:138
static LogicalResult ReduceInferReturnTypes(ShapeAdaptor operandShape, Type inputType, IntegerAttr axis, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
Definition TosaOps.cpp:4005
static void printScaleValues(AsmPrinter &printer, ArrayRef< Attribute > scaleValues, Type)
Definition TosaOps.cpp:547
static void buildAvgPool2dAdaptiveOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, DenseI64ArrayAttr kernel, DenseI64ArrayAttr stride, DenseI64ArrayAttr pad, TypeAttr accType)
This builder mirrors avg_pool2d quant-info handling and materializes kernel/stride/pad as const_shape...
Definition TosaOps.cpp:1435
static LogicalResult verifyRescaleValueAndZpTypes(Operation *op, Value val, Value valZp, StringRef name)
Definition TosaOps.cpp:296
static void printOptionalBoolClause(OpAsmPrinter &printer, StringRef keyword, BoolAttr attr)
Definition TosaOps.cpp:5969
static LogicalResult errorIfShapeNotSizeOne(Operation *op, Type type)
Definition TosaOps.cpp:968
static void printLocalBound(OpAsmPrinter &printer, Operation *, BoolAttr attr)
Definition TosaOps.cpp:5982
LogicalResult argMaxMinVerify(T op)
Definition TosaOps.cpp:625
static LogicalResult verifyMatMulZeroPointType(T op, Value input, Value zp, StringRef inputName, StringRef zpName)
Definition TosaOps.cpp:2036
static ParseResult parseScaleValues(AsmParser &parser, SmallVector< Attribute > &scaleValues, Type scaleType)
Definition TosaOps.cpp:518
static ParseResult parseInputUnsigned(OpAsmParser &parser, BoolAttr &result)
Definition TosaOps.cpp:5986
#define REDUCE_SHAPE_INFER(OP)
Definition TosaOps.cpp:4030
static LogicalResult verifyConvOp(T op)
Definition TosaOps.cpp:656
static LogicalResult verifyAvgPoolCommonTypeAndZpChecks(T op)
Definition TosaOps.cpp:1143
static LogicalResult verifyVariableOpErrorIf(T op, Type type, StringRef name)
Definition TosaOps.cpp:977
static LogicalResult poolingInferReturnTypes(ShapeAdaptor inputShape, ArrayRef< int64_t > kernel, ArrayRef< int64_t > stride, ArrayRef< int64_t > pad, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
Definition TosaOps.cpp:4220
static void buildPadOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, Value paddings)
This builder is called on TOSA pad operator that needs to create its own OptionalAttr quantization_at...
Definition TosaOps.cpp:1522
static LogicalResult verifyPoolingOpImpl(Operation *op, ArrayRef< int64_t > kernel, ArrayRef< int64_t > strides, ArrayRef< int64_t > padding, Value input, Value output)
Definition TosaOps.cpp:1042
static std::optional< int64_t > idivCheck(const int64_t lhs, const int64_t rhs)
Definition TosaOps.cpp:279
static void buildVariableOp(OpBuilder &builder, OperationState &result, StringRef name, Type variableType, Attribute initialValue)
Definition TosaOps.cpp:1536
static void buildMatMulLikeOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value a, Value b)
Definition TosaOps.cpp:1350
LogicalResult verifyConvOutputSize(Operation *op, const int64_t inputSize, const int64_t kernelSize, const int64_t outputSize, const int64_t padBefore, const int64_t padAfter, const int64_t stride, const int64_t dilation, const llvm::StringRef dimName, const llvm::StringRef dimAxis, const llvm::StringRef padBeforeName, const llvm::StringRef padAfterName)
Definition TosaOps.cpp:385
static LogicalResult verifyReduceOp(T op)
Definition TosaOps.cpp:4055
#define NARY_SHAPE_INFER(OP)
Definition TosaOps.cpp:4123
#define ZERO_POINT_HELPER(OP, OPERAND_NAME, SIGN_EXTEND)
Definition TosaOps.cpp:3229
static void buildTransConvOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, Value weight, Value bias, DenseI64ArrayAttr outpad, DenseI64ArrayAttr stride, TypeAttr accType)
Handles tosa.transpose_conv2d which has outpad and output shape attributes.
Definition TosaOps.cpp:1332
LogicalResult inferArgMaxMinReturnTypeComponents(MLIRContext *context, ::std::optional< Location > location, A adaptor, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
Definition TosaOps.cpp:596
static void extractAdaptivePoolingConstShapeOperands(T op, AdaptivePoolingConstShapeValues &values)
Definition TosaOps.cpp:1201
static LogicalResult verifyConvOpErrorIf(T op)
Definition TosaOps.cpp:823
static FailureOr< int64_t > getZeroPoint(Value val, bool signExtend)
Definition TosaOps.cpp:3160
static constexpr bool IsSupportedAdaptivePoolConstShapeVerifyOp
Definition TosaOps.cpp:1194
LogicalResult tryUpdateDimOrFailure(Operation *op, int64_t &currDim, const int64_t newDim, const StringRef operandName, const StringRef dimName)
Definition TosaOps.cpp:341
static LogicalResult verifyConvOpModes(T op)
Definition TosaOps.cpp:801
static LogicalResult NAryInferReturnTypes(const ValueShapeRange &operands, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
Definition TosaOps.cpp:4111
#define COMPATIBLE_RETURN_TYPES(OP)
Definition TosaOps.cpp:4021
static LogicalResult resolveBroadcastShape(const ValueShapeRange &operands, SmallVector< int64_t > &outShape)
Definition TosaOps.cpp:1580
static LogicalResult verifyMatMulQuantizedOperandsType(T op, Type aElementType, Type bElementType)
Definition TosaOps.cpp:2009
static LogicalResult verifyOutputShapeCompatibleWithExpected(Operation *op, ShapedType outputType, ArrayRef< int64_t > expectedShape, StringRef outputName="output")
Definition TosaOps.cpp:369
static void printInputUnsigned(OpAsmPrinter &printer, Operation *, BoolAttr attr)
Definition TosaOps.cpp:5990
static void buildNegateOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input)
This builder is called on single-parameter negate operator to construct input and output zero points ...
Definition TosaOps.cpp:1482
static void buildConvOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, Value weight, Value bias, DenseI64ArrayAttr pad, DenseI64ArrayAttr stride, DenseI64ArrayAttr dilation, TypeAttr accType)
This builder is called on all convolution operators except TransposeConv, which has specialized outpu...
Definition TosaOps.cpp:1308
static SmallVector< int64_t > getMatMulBatchShape(const ShapeAdaptor &shape)
Definition TosaOps.cpp:2061
static void buildAvgPool2dOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, DenseArrayAttr kernel, DenseArrayAttr stride, DenseArrayAttr pad, TypeAttr accType)
Both the tosa.avg_pool2d and unary ops use the same UnaryOpQuantizationAttr but avg_pool operator has...
Definition TosaOps.cpp:1391
static LogicalResult errorIfTypeOrShapeMismatch(Operation *op, Type type1, StringRef name1, Type type2, StringRef name2)
Definition TosaOps.cpp:926
static LogicalResult inferMatMulReturnTypeComponents(const ShapeAdaptor &aShape, const ShapeAdaptor &bShape, bool transposeB, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
Definition TosaOps.cpp:1967
static void buildMatMulTOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value a, Value b)
Definition TosaOps.cpp:1381
static FailureOr< SmallVector< int64_t > > resolveMatMulOutputShape(const ShapeAdaptor &aShape, const ShapeAdaptor &bShape, int64_t outputRank, bool transposeB)
Definition TosaOps.cpp:1940
static FailureOr< int64_t > resolveBroadcastDim(const int64_t dim1, const int64_t dim2)
Definition TosaOps.cpp:1566
static LogicalResult verifyZeroPoint(T op, Value val, const int64_t &zp, const std::string &operand)
Definition TosaOps.cpp:3187
static LogicalResult verifyPoolingOp(T op)
Definition TosaOps.cpp:1137
static LogicalResult verifyDimIsPowerOfTwo(Operation *op, const int64_t dimSize, const llvm::StringRef dimName)
Definition TosaOps.cpp:1657
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Definition Traits.cpp:117
static void updateIfDynamic(int64_t &current, int64_t candidate)
Definition TosaOps.cpp:4259
void inferWeightShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &weightSpatial)
Definition TosaOps.cpp:4352
LogicalResult getSpatialParameters(SmallVector< int64_t > &padValues, SmallVector< int64_t > &strideValues, SmallVector< int64_t > &dilationValues)
Definition TosaOps.cpp:4381
void inferInputShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &inputSpatial)
Definition TosaOps.cpp:4326
ConvInferShapeAdaptor(Conv2DBlockScaledOp::Adaptor adaptor)
Definition TosaOps.cpp:4323
void inferInputShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &inputSpatial)
Definition TosaOps.cpp:4272
void inferWeightShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &weightSpatial)
Definition TosaOps.cpp:4287
ConvInferShapeAdaptor(Conv2DOp::Adaptor adaptor)
Definition TosaOps.cpp:4269
LogicalResult getSpatialParameters(SmallVector< int64_t > &padValues, SmallVector< int64_t > &strideValues, SmallVector< int64_t > &dilationValues)
Definition TosaOps.cpp:4305
void inferWeightShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &weightSpatial)
Definition TosaOps.cpp:4422
ConvInferShapeAdaptor(Conv3DOp::Adaptor adaptor)
Definition TosaOps.cpp:4402
void inferInputShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &inputSpatial)
Definition TosaOps.cpp:4405
LogicalResult getSpatialParameters(SmallVector< int64_t > &padValues, SmallVector< int64_t > &strideValues, SmallVector< int64_t > &dilationValues)
Definition TosaOps.cpp:4442
This base class exposes generic asm parser hooks, usable across the various derived parsers.
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalEqual()=0
Parse a = token if present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseOptionalColon()=0
Parse a : token if present.
virtual ParseResult parseOptionalAttrDictWithKeyword(NamedAttrList &result)=0
Parse a named dictionary into 'result' if the attributes keyword is present.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseColon()=0
Parse a : token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
virtual ParseResult parseFloat(double &result)=0
Parse a floating point value from the stream.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
This base class exposes generic asm printer hooks, usable across the various derived printers.
virtual void printAttributeWithoutType(Attribute attr)
Print the given attribute without its type.
virtual void printAttribute(Attribute attr)
void printArrowTypeList(TypeRange &&types)
Attributes are known-constant values of operations.
Definition Attributes.h:25
MutableArrayRef< BlockArgument > BlockArgListType
Definition Block.h:110
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
IntegerType getI32Type()
Definition Builders.cpp:71
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
StringAttr getStringAttr(const Twine &bytes)
Definition Builders.cpp:271
DenseIntElementsAttr getIndexTensorAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:201
An attribute that represents a reference to a dense vector or tensor object.
auto getValues() const
Return the held element values as a range of the given type.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
An attribute that represents a reference to a dense integer vector or tensor object.
This class contains all of the information necessary to report a diagnostic to the DiagnosticEngine.
virtual InFlightDiagnostic emitError(const Twine &msg={}) const =0
Emit an error to the reader.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
Definition Builders.h:632
This class represents a diagnostic that is inflight and set to be reported.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual OptionalParseResult parseOptionalAssignmentList(SmallVectorImpl< Argument > &lhs, SmallVectorImpl< UnresolvedOperand > &rhs)=0
virtual ParseResult parseRegion(Region &region, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDictWithKeyword(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary prefixed with 'attribute...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
void printFunctionalType(Operation *op)
Print the complete type of an operation in functional form.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
This class helps build Operations.
Definition Builders.h:210
This class indicates that op operates on tosa shape types.
Definition TosaOps.h:78
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
ResultRange result_range
Support result iteration.
Definition Operation.h:435
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
OperandRange operand_range
Definition Operation.h:396
operand_type_range getOperandTypes()
Definition Operation.h:422
result_type_range getResultTypes()
Definition Operation.h:453
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
ParseResult value() const
Access the internal ParseResult value.
bool has_value() const
Returns true if we contain a valid ParseResult value.
Type-safe wrapper around a void* for passing properties, including the properties structs of operatio...
This class provides an abstraction over the different types of ranges over Regions.
Definition Region.h:363
bool empty()
Definition Region.h:60
This diagnostic handler is a simple RAII class that registers and erases a diagnostic handler on a gi...
Adaptor class to abstract the differences between whether value is from a ShapedType or ShapedTypeCom...
bool isDynamicDim(int index) const
Returns whether the index'th dimension is dynamic.
int64_t getDimSize(int index) const
Returns the size of the index'th dimension.
int64_t getRank() const
Returns the rank of the shape.
bool hasStaticShape() const
Returns whether the shape is fully static.
int64_t getNumElements() const
Returns the number of elements in the shape.
void getDims(SmallVectorImpl< int64_t > &res) const
Populates the dimensions from shape referenced.
bool hasRank() const
Returns whether the shape has a rank.
ShapedTypeComponents that represents the components of a ShapedType.
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
Definition SymbolTable.h:25
Operation * lookup(StringRef name) const
Look up a symbol with the specified name, returning null if no such name exists.
Tensor types represent multi-dimensional arrays, and have two variants: RankedTensorType and Unranked...
ArrayRef< int64_t > getShape() const
Returns the shape of this tensor type.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isSignlessInteger() const
Return true if this is a signless integer type (with the specified width).
Definition Types.cpp:66
bool isF32() const
Definition Types.cpp:40
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
Definition Types.cpp:90
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
bool isF16() const
Definition Types.cpp:38
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
bool isBF16() const
Definition Types.cpp:37
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getTypes() const
Range of values and shapes (corresponding effectively to Shapes dialect's ValueShape type concept).
ShapeAdaptor getShape(int index) const
Returns the shape of index'th operand.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
LogicalResult verifyAtLeastNOperands(Operation *op, unsigned numOperands)
LogicalResult verifyTosaShapeOperatorWithSameRanks(Operation *op)
Definition TosaOps.cpp:5808
LogicalResult verifyTosaResolvableShapeOperands(Operation *op)
Definition TosaOps.cpp:5795
bool getBroadcastedShape(ArrayRef< int64_t > shape1, ArrayRef< int64_t > shape2, SmallVectorImpl< int64_t > &resultShape)
Returns true and sets resultShape to the broadcasted shape from the two given shapes if they are broa...
Definition Traits.cpp:59
LogicalResult convertFloatTypeFromAttribute(Type type, Attribute attr, llvm::SmallVectorImpl< char > &result)
Float type implementation of DenseElementTypeInterface::convertFromAttribute.
Attribute convertFloatTypeToAttribute(Type type, llvm::ArrayRef< char > rawData)
Float type implementation of DenseElementTypeInterface::convertToAttribute.
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
Definition Utils.cpp:18
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
SmallVector< unsigned > getBlockSize(AffineMap dimToLvl)
Given the dimToLvl map, returns the block sizes in a vector.
ConvOpQuantizationAttr buildConvOpQuantizationAttr(OpBuilder &builder, Value input, Value weight)
Method to build ConvOpQuantizationAttr, called from ConvOpQuantInfoBuilder/TransConvOpQuantInfoBuilde...
Type getStorageElementTypeOrSelf(Type type)
Definition TosaOps.cpp:285
RankedTensorType getVariableType(VariableOp variableOp)
Type buildConvOpResultTypeInfo(OpBuilder &builder, Type outputType, Value input, Value weight)
construct ConvOp output type with correct bitwidth based on input/weight width.
ParseResult parseVariableOpTypeOrInitialValue(OpAsmParser &parser, DenseElementsAttr &varShapeAttr, TypeAttr &typeAttr, Attribute &initialValueAttr)
Definition TosaOps.cpp:227
PadOpQuantizationAttr buildPadOpQuantizationAttr(OpBuilder &builder, Value input)
Builds PadOpQuantizationAttr, called from PadOpQuantInfoBuilder: inputZp: input zeropoint.
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
std::pair< Value, Value > createZPsAsConst(OpBuilder &builder, Value input, Value weight)
void printVariableOpTypeOrInitialValue(OpAsmPrinter &p, Operation *op, DenseElementsAttr varShapeAttr, TypeAttr typeAttr, Attribute initialValueAttr)
Definition TosaOps.cpp:252
FailureOr< T > getConstantScalarIntValue(Value val)
Value getTosaConstShape(ImplicitLocOpBuilder &builder, llvm::ArrayRef< int64_t > shape)
MatMulOpQuantizationAttr buildMatMulOpQuantizationAttr(OpBuilder &builder, Value a, Value b)
Builds MatMulOpQuantizationAttr, called from MatMulOpQuantInfoBuilder: aZp: input a zeropoint bZp: in...
unsigned getBitWidth(Type type)
Definition TosaOps.cpp:331
std::optional< Value > createZeroPointTensor(OpBuilder &builder, Location loc, Type srcElemType, int64_t zp=0)
Definition TosaOps.cpp:5759
bool isa_tosa_shape_type(mlir::Type t)
Definition TosaOps.cpp:5783
SmallVector< int64_t > convertFromMlirShape(ArrayRef< int64_t > shape)
UnaryOpQuantizationAttr buildUnaryOpQuantizationAttr(OpBuilder &builder, Value input, Type outputRawType)
Builds UnaryOpQuantizationAttr UnaryOpQuantInfoBuilder: inputZp: input zeropoint outputZp: output zer...
Type getStorageElementTypeFromQuantized(quant::QuantizedType quantizedType)
Value createPadConstTensor(OpBuilder &builder, Location loc, Value src, int32_t val=0)
Definition TosaOps.cpp:316
LogicalResult verifyBlockScaledTensorType(mlir::Type type, llvm::function_ref< mlir::InFlightDiagnostic()> emitError=nullptr, bool allowScaleValues=false)
Definition TosaOps.cpp:444
std::string getTosaTensorTypeErrorMessage(mlir::Type type)
Definition TosaOps.cpp:503
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
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
LogicalResult verifyCompatibleShapes(TypeRange types1, TypeRange types2)
Returns success if the given two arrays have the same number of elements and each pair wise entries h...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
LogicalResult emitOptionalError(std::optional< Location > loc, Args &&...args)
Overloads of the above emission functions that take an optionally null location.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
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.
LogicalResult verifyRanksMatch(Operation *op, ShapedType lhs, ShapedType rhs, StringRef lhsName, StringRef rhsName)
Verify that two shaped types have matching ranks.
LogicalResult verifyCompatibleShape(ArrayRef< int64_t > shape1, ArrayRef< int64_t > shape2)
Returns success if the given two shapes are compatible.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
bool isPermutationVector(ArrayRef< int64_t > interchange)
Method to check if an interchange vector is a permutation.
This represents an operation in an abstracted form, suitable for use with the builder APIs.
static ValueKnowledge meet(const ValueKnowledge &lhs, const ValueKnowledge &rhs)
Definition ShapeUtils.h:136
static ValueKnowledge getKnowledgeFromType(Type type)
Definition ShapeUtils.h:45