MLIR 24.0.0git
TosaToLinalg.cpp
Go to the documentation of this file.
1//===- TosaToLinalg.cpp - Lowering Tosa to Linalg Dialect -----------------===//
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// These rewriters lower from the Tosa to the Linalg dialect.
10//
11//===----------------------------------------------------------------------===//
12
25#include "mlir/IR/Matchers.h"
29#include "llvm/ADT/STLExtras.h"
30#include "llvm/ADT/Sequence.h"
31#include "llvm/ADT/SmallVectorExtras.h"
32
33#include <type_traits>
34
35using namespace mlir;
36using namespace mlir::tosa;
37
38template <typename OpTy>
40 TypeRange resultTypes,
41 ValueRange operands) {
42 typename OpTy::Properties properties{};
43 OpTy::populateDefaultProperties(
44 OperationName(OpTy::getOperationName(), builder.getContext()),
45 properties);
46 return OpTy::create(builder, loc, resultTypes, operands, properties,
47 /*discardableAttributes=*/{});
48}
49
50// Helper function to materialize the semantically correct compare and select
51// operations given a binary operation with a specific NaN propagation mode.
52//
53// In the case of "PROPAGATE" semantics no compare and selection is required and
54// this function does nothing.
55//
56// In the case of "IGNORE" semantics this function materializes a comparison of
57// the current operands to the op which will return true for any NaN
58// argument and then selects between the non-NaN operation argument and the
59// calculated result based on whether the lhs or rhs is NaN or not. In pseudo
60// code:
61//
62// In the case that the op is operating on non floating point types we ignore
63// the attribute completely, this is consistent with the TOSA spec which has
64// the following wording: "This attribute is ignored by non floating-point
65// types."
66//
67// binary<op>(lhs, rhs):
68// result = op(lhs, rhs)
69// if lhs == NaN return rhs
70// if rhs == NaN return lhs
71// return result
72template <typename OpTy>
73static Value
75 Value lhs, Value rhs, Value result) {
76 // NaN propagation has no meaning for non floating point types.
77 if (!isa<FloatType>(getElementTypeOrSelf(lhs)))
78 return result;
79
80 auto nanMode = op.getNanMode();
81 if (nanMode == NanPropagationMode::PROPAGATE)
82 return result;
83
84 // Unordered comparison of NaN against itself will always return true.
85 Value lhsIsNaN = arith::CmpFOp::create(rewriter, op.getLoc(),
86 arith::CmpFPredicate::UNO, lhs, lhs);
87 Value rhsIsNaN = arith::CmpFOp::create(rewriter, op.getLoc(),
88 arith::CmpFPredicate::UNO, rhs, rhs);
89 Value rhsOrResult =
90 arith::SelectOp::create(rewriter, op.getLoc(), lhsIsNaN, rhs, result);
91 return arith::SelectOp::create(rewriter, op.getLoc(), rhsIsNaN, lhs,
92 rhsOrResult);
93}
94
96 Operation *op, ValueRange args, ArrayRef<Type> resultTypes,
97 ConversionPatternRewriter &rewriter) {
98 Location loc = op->getLoc();
99 auto elementTy =
100 cast<ShapedType>(op->getOperand(0).getType()).getElementType();
101
102 // tosa::AbsOp
103 if (isa<tosa::AbsOp>(op) && isa<FloatType>(elementTy))
104 return createWithDefaultProperties<math::AbsFOp>(rewriter, loc, resultTypes,
105 args);
106
107 if (isa<tosa::AbsOp>(op) && isa<IntegerType>(elementTy)) {
108 auto zero = arith::ConstantOp::create(rewriter, loc,
109 rewriter.getZeroAttr(elementTy));
110 auto neg = arith::SubIOp::create(rewriter, loc, zero, args[0]);
111 return arith::MaxSIOp::create(rewriter, loc, args[0], neg);
112 }
113
114 // tosa::AddOp
115 if (isa<tosa::AddOp>(op) && isa<FloatType>(elementTy))
117 resultTypes, args);
118
119 if (isa<tosa::AddOp>(op) && isa<IntegerType>(elementTy))
121 resultTypes, args);
122
123 // tosa::SubOp
124 if (isa<tosa::SubOp>(op) && isa<FloatType>(elementTy))
126 resultTypes, args);
127
128 if (isa<tosa::SubOp>(op) && isa<IntegerType>(elementTy))
130 resultTypes, args);
131
132 // tosa::IntDivOp
133 if (isa<tosa::IntDivOp>(op) && isa<IntegerType>(elementTy))
135 resultTypes, args);
136
137 // tosa::ReciprocalOp
138 if (isa<tosa::ReciprocalOp>(op) && isa<FloatType>(elementTy)) {
139 auto one =
140 arith::ConstantOp::create(rewriter, loc, FloatAttr::get(elementTy, 1));
141 return arith::DivFOp::create(rewriter, loc, one, args[0]);
142 }
143
144 // tosa::MulOp
145 if (isa<tosa::MulOp>(op)) {
146 auto shiftVal = cast<tosa::MulOp>(op).getShift();
147 DenseElementsAttr shiftElem;
148 bool shiftIsConstant = true;
149 int32_t shift = 0;
150 if (matchPattern(shiftVal, m_Constant(&shiftElem)))
151 shift = shiftElem.getValues<IntegerAttr>()[0].getInt();
152 else
153 shiftIsConstant = false;
154
155 if (isa<FloatType>(elementTy)) {
156 if (shift != 0) {
157 (void)rewriter.notifyMatchFailure(op,
158 "Cannot have shift value for float");
159 return nullptr;
160 }
161 return arith::MulFOp::create(rewriter, loc, args[0], args[1]);
162 }
163
164 if (isa<IntegerType>(elementTy)) {
165 Value a = args[0];
166 Value b = args[1];
167
168 if (shift > 0 || !shiftIsConstant) {
169 Value shiftConst;
170 if (shiftIsConstant)
171 shiftConst = arith::ConstantIntOp::create(rewriter, loc, shift,
172 /*bitwidth=*/8);
173
174 if (!a.getType().isInteger(32))
175 a = arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), a);
176
177 if (!b.getType().isInteger(32))
178 b = arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), b);
179
180 auto shiftAmount = shiftIsConstant ? shiftConst : args[2];
181 auto roundingAttr = RoundingModeAttr::get(rewriter.getContext(),
182 RoundingMode::SINGLE_ROUND);
183 auto result =
184 tosa::ApplyScaleOp::create(rewriter, loc, rewriter.getI32Type(), a,
185 b, shiftAmount, roundingAttr);
186
187 return result;
188 }
189
190 int aWidth = a.getType().getIntOrFloatBitWidth();
191 int bWidth = b.getType().getIntOrFloatBitWidth();
192 int cWidth = resultTypes[0].getIntOrFloatBitWidth();
193
194 if (aWidth < cWidth)
195 a = arith::ExtSIOp::create(rewriter, loc, resultTypes[0], a);
196 if (bWidth < cWidth)
197 b = arith::ExtSIOp::create(rewriter, loc, resultTypes[0], b);
198
199 return arith::MulIOp::create(rewriter, loc, resultTypes, a, b);
200 }
201 }
202
203 // tosa::NegateOp
204 if (isa<tosa::NegateOp>(op)) {
205 auto negate = cast<tosa::NegateOp>(op);
206
207 int64_t inZp = 0, outZp = 0;
208 FailureOr<int64_t> maybeInZp = negate.getInput1ZeroPoint();
209 FailureOr<int64_t> maybeOutZp = negate.getOutputZeroPoint();
210 bool hasInZp = !failed(maybeInZp);
211 bool hasOutZp = !failed(maybeOutZp);
212 if (hasInZp)
213 inZp = *maybeInZp;
214 if (hasOutZp)
215 outZp = *maybeOutZp;
216
217 if (isa<FloatType>(elementTy))
218 return arith::NegFOp::create(rewriter, loc, resultTypes, args[0]);
219
220 if (isa<IntegerType>(elementTy)) {
221 Value zpAddValue;
222 Type intermediateType;
223 // Compute the maximum value that can occur in the intermediate buffer.
224 const int32_t inputBitWidth = elementTy.getIntOrFloatBitWidth();
225 int intermediateBitWidth = 64;
226
227 if (hasInZp && hasOutZp) {
228 // Compute the maximum value that can occur in the intermediate buffer.
229 const int64_t zpAdd = inZp + outZp;
230 const int64_t maxValue =
231 APInt::getSignedMaxValue(inputBitWidth).getSExtValue() +
232 std::abs(zpAdd) + 1;
233
234 // Convert that maximum value into the maximum bitwidth needed to
235 // represent it.
236 if (maxValue <= APInt::getSignedMaxValue(16).getSExtValue()) {
237 intermediateBitWidth = 16;
238 } else if (maxValue <= APInt::getSignedMaxValue(32).getSExtValue()) {
239 intermediateBitWidth = 32;
240 }
241
242 intermediateType = rewriter.getIntegerType(intermediateBitWidth);
243 zpAddValue = arith::ConstantOp::create(
244 rewriter, loc, rewriter.getIntegerAttr(intermediateType, zpAdd));
245 } else {
246 intermediateType = rewriter.getIntegerType(intermediateBitWidth);
247 Value arg1 = args[1];
248 Value arg2 = args[2];
249 // Avoid verifier-invalid no-op sign-extends; only widen when needed.
250 if (arg1.getType() != intermediateType)
251 arg1 = arith::ExtSIOp::create(rewriter, loc, intermediateType, arg1);
252 if (arg2.getType() != intermediateType)
253 arg2 = arith::ExtSIOp::create(rewriter, loc, intermediateType, arg2);
254 zpAddValue =
255 arith::AddIOp::create(rewriter, loc, intermediateType, arg1, arg2);
256 }
257
258 // The negation can be applied by doing:
259 // outputValue = inZp + outZp - inputValue
260 Value ext = args[0];
261 if (ext.getType() != intermediateType)
262 ext = arith::ExtSIOp::create(rewriter, loc, intermediateType, ext);
263 auto sub = arith::SubIOp::create(rewriter, loc, zpAddValue, ext);
264
265 // Clamp to the negation range.
267 rewriter, loc, intermediateType,
268 APInt::getSignedMinValue(inputBitWidth).getSExtValue());
270 rewriter, loc, intermediateType,
271 APInt::getSignedMaxValue(inputBitWidth).getSExtValue());
272 auto clamp = clampIntHelper(loc, sub, min, max, rewriter, false);
273
274 // Truncate to the final value, skipping no-op trunci when widths match.
275 if (clamp.getType() == elementTy)
276 return clamp;
277 return arith::TruncIOp::create(rewriter, loc, elementTy, clamp);
278 }
279 }
280
281 // tosa::BitwiseAndOp
282 if (isa<tosa::BitwiseAndOp>(op) && isa<IntegerType>(elementTy))
283 return arith::AndIOp::create(rewriter, loc, resultTypes, args);
284
285 // tosa::BitwiseOrOp
286 if (isa<tosa::BitwiseOrOp>(op) && isa<IntegerType>(elementTy))
287 return arith::OrIOp::create(rewriter, loc, resultTypes, args);
288
289 // tosa::BitwiseNotOp
290 if (isa<tosa::BitwiseNotOp>(op) && isa<IntegerType>(elementTy)) {
291 auto allOnesAttr = rewriter.getIntegerAttr(
292 elementTy, APInt::getAllOnes(elementTy.getIntOrFloatBitWidth()));
293 auto allOnes = arith::ConstantOp::create(rewriter, loc, allOnesAttr);
294 return arith::XOrIOp::create(rewriter, loc, resultTypes, args[0], allOnes);
295 }
296
297 // tosa::BitwiseXOrOp
298 if (isa<tosa::BitwiseXorOp>(op) && isa<IntegerType>(elementTy))
299 return arith::XOrIOp::create(rewriter, loc, resultTypes, args);
300
301 // tosa::LogicalLeftShiftOp
302 if (isa<tosa::LogicalLeftShiftOp>(op) && isa<IntegerType>(elementTy))
304 resultTypes, args);
305
306 // tosa::LogicalRightShiftOp
307 if (isa<tosa::LogicalRightShiftOp>(op) && isa<IntegerType>(elementTy))
309 resultTypes, args);
310
311 // tosa::ArithmeticRightShiftOp
312 if (isa<tosa::ArithmeticRightShiftOp>(op) && isa<IntegerType>(elementTy)) {
314 rewriter, loc, resultTypes, args);
315 bool round = cast<tosa::ArithmeticRightShiftOp>(op).getRound();
316 if (!round) {
317 return result;
318 }
319
320 Type i1Ty = IntegerType::get(rewriter.getContext(), /*width=*/1);
321 auto one = arith::ConstantOp::create(rewriter, loc,
322 IntegerAttr::get(elementTy, 1));
323 auto zero = arith::ConstantOp::create(rewriter, loc,
324 IntegerAttr::get(elementTy, 0));
325 auto i1zero =
326 arith::ConstantOp::create(rewriter, loc, IntegerAttr::get(i1Ty, 0));
327 auto i1one =
328 arith::ConstantOp::create(rewriter, loc, IntegerAttr::get(i1Ty, 1));
329
330 // Checking that input2 != 0
331 auto shiftValueGreaterThanZero = arith::CmpIOp::create(
332 rewriter, loc, arith::CmpIPredicate::sgt, args[1], zero);
333
334 // Checking for the last bit of input1 to be 1
335 auto subtract =
336 arith::SubIOp::create(rewriter, loc, resultTypes, args[1], one);
337 auto shifted =
338 arith::ShRSIOp::create(rewriter, loc, resultTypes, args[0], subtract)
339 ->getResults();
341 rewriter, loc, TypeRange{i1Ty}, shifted);
342 auto isInputOdd =
343 arith::AndIOp::create(rewriter, loc, i1Ty, truncated, i1one);
344 // shifted, truncated, isInputOdd can be poison when input2 is 0.
345 auto shouldRound = arith::SelectOp::create(
346 rewriter, loc, i1Ty, shiftValueGreaterThanZero, isInputOdd, i1zero);
347 auto extended =
348 arith::ExtUIOp::create(rewriter, loc, resultTypes, shouldRound);
349 return arith::AddIOp::create(rewriter, loc, resultTypes, result, extended);
350 }
351
352 // tosa::ClzOp
353 if (isa<tosa::ClzOp>(op) && isa<IntegerType>(elementTy)) {
354 return math::CountLeadingZerosOp::create(rewriter, loc, elementTy, args[0]);
355 }
356
357 // tosa::LogicalAnd
358 if (isa<tosa::LogicalAndOp>(op) && elementTy.isInteger(1))
359 return arith::AndIOp::create(rewriter, loc, resultTypes, args);
360
361 // tosa::LogicalNot
362 if (isa<tosa::LogicalNotOp>(op) && elementTy.isInteger(1)) {
363 auto one = arith::ConstantOp::create(rewriter, loc,
364 rewriter.getIntegerAttr(elementTy, 1));
365 return arith::XOrIOp::create(rewriter, loc, resultTypes, args[0], one);
366 }
367
368 // tosa::LogicalOr
369 if (isa<tosa::LogicalOrOp>(op) && elementTy.isInteger(1))
370 return arith::OrIOp::create(rewriter, loc, resultTypes, args);
371
372 // tosa::LogicalXor
373 if (isa<tosa::LogicalXorOp>(op) && elementTy.isInteger(1))
374 return arith::XOrIOp::create(rewriter, loc, resultTypes, args);
375
376 // tosa::PowOp
377 if (isa<tosa::PowOp>(op) && isa<FloatType>(elementTy))
379 resultTypes, args);
380
381 // tosa::RsqrtOp
382 if (isa<tosa::RsqrtOp>(op) && isa<FloatType>(elementTy))
384 resultTypes, args);
385
386 // tosa::LogOp
387 if (isa<tosa::LogOp>(op) && isa<FloatType>(elementTy))
389 resultTypes, args);
390
391 // tosa::ExpOp
392 if (isa<tosa::ExpOp>(op) && isa<FloatType>(elementTy))
394 resultTypes, args);
395
396 // tosa::SinOp
397 if (isa<tosa::SinOp>(op) && isa<FloatType>(elementTy))
399 resultTypes, args);
400
401 // tosa::CosOp
402 if (isa<tosa::CosOp>(op) && isa<FloatType>(elementTy))
404 resultTypes, args);
405
406 // tosa::TanhOp
407 if (isa<tosa::TanhOp>(op) && isa<FloatType>(elementTy))
409 resultTypes, args);
410
411 // tosa::ErfOp
412 if (isa<tosa::ErfOp>(op) && llvm::isa<FloatType>(elementTy))
414 resultTypes, args);
415
416 // tosa::GreaterOp
417 if (isa<tosa::GreaterOp>(op) && isa<FloatType>(elementTy))
418 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OGT,
419 args[0], args[1]);
420
421 if (isa<tosa::GreaterOp>(op) && elementTy.isSignlessInteger())
422 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::sgt,
423 args[0], args[1]);
424
425 // tosa::GreaterEqualOp
426 if (isa<tosa::GreaterEqualOp>(op) && isa<FloatType>(elementTy))
427 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OGE,
428 args[0], args[1]);
429
430 if (isa<tosa::GreaterEqualOp>(op) && elementTy.isSignlessInteger())
431 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::sge,
432 args[0], args[1]);
433
434 // tosa::EqualOp
435 if (isa<tosa::EqualOp>(op) && isa<FloatType>(elementTy))
436 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OEQ,
437 args[0], args[1]);
438
439 if (isa<tosa::EqualOp>(op) && elementTy.isSignlessInteger())
440 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq,
441 args[0], args[1]);
442
443 // tosa::SelectOp
444 if (isa<tosa::SelectOp>(op)) {
445 elementTy = cast<ShapedType>(op->getOperand(1).getType()).getElementType();
446 if (isa<FloatType>(elementTy) || isa<IntegerType>(elementTy))
447 return arith::SelectOp::create(rewriter, loc, args[0], args[1], args[2]);
448 }
449
450 // tosa::MaximumOp
451 if (isa<tosa::MaximumOp>(op) && isa<FloatType>(elementTy)) {
452 auto max = arith::MaximumFOp::create(rewriter, loc, args[0], args[1]);
453 return materializeBinaryNanCheckIfRequired(llvm::cast<tosa::MaximumOp>(op),
454 rewriter, args[0], args[1], max);
455 }
456
457 if (isa<tosa::MaximumOp>(op) && elementTy.isSignlessInteger()) {
458 return arith::MaxSIOp::create(rewriter, loc, args[0], args[1]);
459 }
460
461 // tosa::MinimumOp
462 if (isa<tosa::MinimumOp>(op) && isa<FloatType>(elementTy)) {
463 auto min = arith::MinimumFOp::create(rewriter, loc, args[0], args[1]);
464 return materializeBinaryNanCheckIfRequired(llvm::cast<tosa::MinimumOp>(op),
465 rewriter, args[0], args[1], min);
466 }
467
468 if (isa<tosa::MinimumOp>(op) && elementTy.isSignlessInteger()) {
469 return arith::MinSIOp::create(rewriter, loc, args[0], args[1]);
470 }
471
472 // tosa::CeilOp
473 if (isa<tosa::CeilOp>(op) && isa<FloatType>(elementTy))
474 return createWithDefaultProperties<math::CeilOp>(rewriter, loc, resultTypes,
475 args);
476
477 // tosa::FloorOp
478 if (isa<tosa::FloorOp>(op) && isa<FloatType>(elementTy))
480 resultTypes, args);
481
482 // tosa::ClampOp
483 if (isa<tosa::ClampOp>(op) && isa<FloatType>(elementTy)) {
484 bool losesInfo = false;
485 auto clampOp = cast<tosa::ClampOp>(op);
486 APFloat minApf = cast<FloatAttr>(clampOp.getMinValAttr()).getValue();
487 APFloat maxApf = cast<FloatAttr>(clampOp.getMaxValAttr()).getValue();
488 minApf.convert(cast<FloatType>(elementTy).getFloatSemantics(),
489 APFloat::rmNearestTiesToEven, &losesInfo);
490 maxApf.convert(cast<FloatType>(elementTy).getFloatSemantics(),
491 APFloat::rmNearestTiesToEven, &losesInfo);
492 auto min = arith::ConstantOp::create(
493 rewriter, loc, elementTy, rewriter.getFloatAttr(elementTy, minApf));
494 auto max = arith::ConstantOp::create(
495 rewriter, loc, elementTy, rewriter.getFloatAttr(elementTy, maxApf));
496 auto result = clampFloatHelper(loc, args[0], min, max, rewriter);
497
498 const auto nanMode = clampOp.getNanMode();
499
500 // NaN propagation has no meaning for non floating point types.
501 if (!isa<FloatType>(elementTy))
502 return result;
503
504 // In the case of "PROPAGATE" semantics no compare and selection is
505 // required.
506 if (nanMode == NanPropagationMode::PROPAGATE)
507 return result;
508
509 // In the case of "IGNORE" semantics materialize a comparison
510 // of the current operand to the reduction which will return true for a NaN
511 // argument and then selects between the initial reduction value and the
512 // calculated result based on whether the argument is NaN or not. In pseudo
513 // code:
514 //
515 // reduce<op>(x, init):
516 // result = op(init, x)
517 // return init if x == NaN else result
518
519 // Unordered comparison of NaN against itself will always return true.
520 Value isNaN = arith::CmpFOp::create(
521 rewriter, op->getLoc(), arith::CmpFPredicate::UNO, args[0], args[0]);
522 // TOSA specifies that in "ignore" NaN mode the result is "min" if the input
523 // is NaN.
524 return arith::SelectOp::create(rewriter, op->getLoc(), isNaN, min, result);
525 }
526
527 if (isa<tosa::ClampOp>(op) && isa<IntegerType>(elementTy)) {
528 auto intTy = cast<IntegerType>(elementTy);
529 auto clampOp = cast<tosa::ClampOp>(op);
530 int64_t min =
531 cast<IntegerAttr>(clampOp.getMinValAttr()).getValue().getSExtValue();
532 int64_t max =
533 cast<IntegerAttr>(clampOp.getMaxValAttr()).getValue().getSExtValue();
534
535 int64_t minRepresentable = std::numeric_limits<int64_t>::min();
536 int64_t maxRepresentable = std::numeric_limits<int64_t>::max();
537 if (intTy.isUnsignedInteger()) {
538 minRepresentable = 0;
539 if (intTy.getIntOrFloatBitWidth() <= 63) {
540 maxRepresentable =
541 (int64_t)APInt::getMaxValue(intTy.getIntOrFloatBitWidth())
542 .getZExtValue();
543 }
544 } else if (intTy.getIntOrFloatBitWidth() <= 64) {
545 // Ensure that min & max fit into signed n-bit constants.
546 minRepresentable = APInt::getSignedMinValue(intTy.getIntOrFloatBitWidth())
547 .getSExtValue();
548 maxRepresentable = APInt::getSignedMaxValue(intTy.getIntOrFloatBitWidth())
549 .getSExtValue();
550 }
551 // Ensure that the bounds are representable as n-bit signed/unsigned
552 // integers.
553 min = std::max(min, minRepresentable);
554 max = std::max(max, minRepresentable);
555 min = std::min(min, maxRepresentable);
556 max = std::min(max, maxRepresentable);
557
558 auto minVal = arith::ConstantIntOp::create(rewriter, loc, min,
559 intTy.getIntOrFloatBitWidth());
560 auto maxVal = arith::ConstantIntOp::create(rewriter, loc, max,
561 intTy.getIntOrFloatBitWidth());
562 return clampIntHelper(loc, args[0], minVal, maxVal, rewriter,
563 intTy.isUnsignedInteger());
564 }
565
566 // tosa::SigmoidOp
567 if (isa<tosa::SigmoidOp>(op) && isa<FloatType>(elementTy)) {
568 auto one =
569 arith::ConstantOp::create(rewriter, loc, FloatAttr::get(elementTy, 1));
570 auto negate = arith::NegFOp::create(rewriter, loc, resultTypes, args[0]);
571 auto exp = mlir::math::ExpOp::create(rewriter, loc, resultTypes, negate);
572 auto added = arith::AddFOp::create(rewriter, loc, exp, one);
573 return arith::DivFOp::create(rewriter, loc, one, added);
574 }
575
576 // tosa::CastOp
577 if (isa<tosa::CastOp>(op)) {
578 Type srcTy = elementTy;
579 Type dstTy = resultTypes.front();
580 if (!srcTy.isIntOrFloat() || !dstTy.isIntOrFloat()) {
581 (void)rewriter.notifyMatchFailure(op, "unsupported type");
582 return nullptr;
583 }
584
585 bool bitExtend =
587
588 if (srcTy == dstTy)
589 return args.front();
590
591 if (isa<FloatType>(srcTy) && isa<FloatType>(dstTy) && bitExtend)
593 resultTypes, args);
594
595 if (isa<FloatType>(srcTy) && isa<FloatType>(dstTy) && !bitExtend)
597 resultTypes, args);
598
599 // 1-bit integers need to be treated as signless.
600 if (srcTy.isInteger(1) && arith::UIToFPOp::areCastCompatible(srcTy, dstTy))
602 resultTypes, args);
603
604 if (srcTy.isInteger(1) && isa<IntegerType>(dstTy) && bitExtend)
606 resultTypes, args);
607
608 // Unsigned integers need an unrealized cast so that they can be passed
609 // to UIToFP.
610 if (srcTy.isUnsignedInteger() && isa<FloatType>(dstTy)) {
611 auto unrealizedCast =
612 UnrealizedConversionCastOp::create(
613 rewriter, loc,
614 rewriter.getIntegerType(srcTy.getIntOrFloatBitWidth()), args[0])
615 .getResult(0);
616 return arith::UIToFPOp::create(rewriter, loc, resultTypes[0],
617 unrealizedCast);
618 }
619
620 // All other si-to-fp conversions should be handled by SIToFP.
621 if (arith::SIToFPOp::areCastCompatible(srcTy, dstTy))
622 return arith::SIToFPOp::create(rewriter, loc, resultTypes, args,
624
625 // Casting to boolean, floats need to only be checked as not-equal to zero.
626 if (isa<FloatType>(srcTy) && dstTy.isInteger(1)) {
627 Value zero = arith::ConstantOp::create(rewriter, loc,
628 rewriter.getFloatAttr(srcTy, 0.0));
629 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::UNE,
630 args.front(), zero);
631 }
632
633 if (arith::FPToSIOp::areCastCompatible(srcTy, dstTy)) {
634 auto rounded = math::RoundEvenOp::create(rewriter, loc, args[0]);
635
636 const auto &fltSemantics = cast<FloatType>(srcTy).getFloatSemantics();
637 // Check whether neither int min nor int max can be represented in the
638 // input floating-point type due to too short exponent range.
639 if (static_cast<int>(dstTy.getIntOrFloatBitWidth()) - 1 >
640 APFloat::semanticsMaxExponent(fltSemantics)) {
641 // Use cmp + select to replace infinites by int min / int max. Other
642 // integral values can be represented in the integer space.
643 auto conv = arith::FPToSIOp::create(rewriter, loc, dstTy, rounded);
644 auto posInf = arith::ConstantOp::create(
645 rewriter, loc,
646 rewriter.getFloatAttr(getElementTypeOrSelf(srcTy),
647 APFloat::getInf(fltSemantics)));
648 auto negInf = arith::ConstantOp::create(
649 rewriter, loc,
650 rewriter.getFloatAttr(
652 APFloat::getInf(fltSemantics, /*Negative=*/true)));
653 auto overflow = arith::CmpFOp::create(
654 rewriter, loc, arith::CmpFPredicate::UEQ, rounded, posInf);
655 auto underflow = arith::CmpFOp::create(
656 rewriter, loc, arith::CmpFPredicate::UEQ, rounded, negInf);
657 auto intMin = arith::ConstantOp::create(
658 rewriter, loc,
659 rewriter.getIntegerAttr(
661 APInt::getSignedMinValue(dstTy.getIntOrFloatBitWidth())));
662 auto intMax = arith::ConstantOp::create(
663 rewriter, loc,
664 rewriter.getIntegerAttr(
666 APInt::getSignedMaxValue(dstTy.getIntOrFloatBitWidth())));
667 auto maxClamped =
668 arith::SelectOp::create(rewriter, loc, overflow, intMax, conv);
669 return arith::SelectOp::create(rewriter, loc, underflow, intMin,
670 maxClamped);
671 }
672
673 auto intMinFP = arith::ConstantOp::create(
674 rewriter, loc,
675 rewriter.getFloatAttr(
677 APInt::getSignedMinValue(dstTy.getIntOrFloatBitWidth())
678 .getSExtValue()));
679
680 // Check whether the mantissa has enough bits to represent int max.
681 if (cast<FloatType>(srcTy).getFPMantissaWidth() >=
682 dstTy.getIntOrFloatBitWidth() - 1) {
683 // Int min can also be represented since it is a power of two and thus
684 // consists of a single leading bit. Therefore we can clamp the input
685 // in the floating-point domain.
686
687 auto intMaxFP = arith::ConstantOp::create(
688 rewriter, loc,
689 rewriter.getFloatAttr(
691 APInt::getSignedMaxValue(dstTy.getIntOrFloatBitWidth())
692 .getSExtValue()));
693
694 Value clamped =
695 clampFloatHelper(loc, rounded, intMinFP, intMaxFP, rewriter);
696 return arith::FPToSIOp::create(rewriter, loc, dstTy, clamped);
697 }
698
699 // Due to earlier check we know exponant range is big enough to represent
700 // int min. We can therefore rely on int max + 1 being representable as
701 // well because it's just int min with a positive sign. So clamp the min
702 // value and compare against that to select the max int value if needed.
703 auto intMaxPlusOneFP = arith::ConstantOp::create(
704 rewriter, loc,
705 rewriter.getFloatAttr(
707 static_cast<double>(
708 APInt::getSignedMaxValue(dstTy.getIntOrFloatBitWidth())
709 .getSExtValue()) +
710 1.0f));
711
712 auto intMax = arith::ConstantOp::create(
713 rewriter, loc,
714 rewriter.getIntegerAttr(
716 APInt::getSignedMaxValue(dstTy.getIntOrFloatBitWidth())));
717 auto minClampedFP =
718 arith::MaximumFOp::create(rewriter, loc, rounded, intMinFP);
719 auto minClamped =
720 arith::FPToSIOp::create(rewriter, loc, dstTy, minClampedFP);
721 auto overflow = arith::CmpFOp::create(
722 rewriter, loc, arith::CmpFPredicate::UGE, rounded, intMaxPlusOneFP);
723 return arith::SelectOp::create(rewriter, loc, overflow, intMax,
724 minClamped);
725 }
726
727 // Casting to boolean, integers need to only be checked as not-equal to
728 // zero.
729 if (isa<IntegerType>(srcTy) && dstTy.isInteger(1)) {
730 Value zero = arith::ConstantIntOp::create(rewriter, loc, 0,
731 srcTy.getIntOrFloatBitWidth());
732 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ne,
733 args.front(), zero);
734 }
735
736 if (isa<IntegerType>(srcTy) && isa<IntegerType>(dstTy) && bitExtend)
737 return arith::ExtSIOp::create(rewriter, loc, resultTypes, args,
739
740 if (isa<IntegerType>(srcTy) && isa<IntegerType>(dstTy) && !bitExtend) {
741 return arith::TruncIOp::create(rewriter, loc, dstTy, args[0]);
742 }
743 }
744
745 (void)rewriter.notifyMatchFailure(
746 op, "unhandled op for linalg body calculation for elementwise op");
747 return nullptr;
748}
749
751
752// Emit an 'arith.constant' op for the given index if it has not been created
753// yet, or return an existing constant. This will prevent an excessive creation
754// of redundant constants, easing readability of emitted code for unit tests.
756 IndexPool &indexPool, int64_t index) {
757 auto [it, inserted] = indexPool.try_emplace(index);
758 if (inserted)
759 it->second =
760 arith::ConstantOp::create(rewriter, loc, rewriter.getIndexAttr(index));
761 return it->second;
762}
763
765 IndexPool &indexPool, Value tensor, int64_t index) {
766 auto indexValue = createIndex(rewriter, loc, indexPool, index);
767 return tensor::DimOp::create(rewriter, loc, tensor, indexValue).getResult();
768}
769
771 IndexPool &indexPool, Value tensor,
772 int64_t index) {
773 auto shapedType = dyn_cast<ShapedType>(tensor.getType());
774 assert(shapedType && shapedType.hasRank() && "expected a ranked shaped type");
775 assert(index >= 0 && index < shapedType.getRank() && "index out of bounds");
776 if (shapedType.isDynamicDim(index))
777 return getTensorDim(rewriter, loc, indexPool, tensor, index);
778 return rewriter.getIndexAttr(shapedType.getDimSize(index));
779}
780
781static bool operandsAndResultsRanked(Operation *operation) {
782 auto isRanked = [](Value value) {
783 return isa<RankedTensorType>(value.getType());
784 };
785 return llvm::all_of(operation->getOperands(), isRanked) &&
786 llvm::all_of(operation->getResults(), isRanked);
787}
788
789// Compute the runtime dimension size for dimension 'dim' of the output by
790// inspecting input 'operands', all of which are expected to have the same rank.
791// This function returns a pair {targetSize, masterOperand}.
792//
793// The runtime size of the output dimension is returned either as a statically
794// computed attribute or as a runtime SSA value.
795//
796// If the target size was inferred directly from one dominating operand, that
797// operand is returned in 'masterOperand'. If the target size is inferred from
798// multiple operands, 'masterOperand' is set to nullptr.
799static std::pair<OpFoldResult, Value>
801 ValueRange operands, int64_t dim) {
802 // If any input operand contains a static size greater than 1 for this
803 // dimension, that is the target size. An occurrence of an additional static
804 // dimension greater than 1 with a different value is undefined behavior.
805 for (auto operand : operands) {
806 auto size = cast<RankedTensorType>(operand.getType()).getDimSize(dim);
807 if (ShapedType::isStatic(size) && size > 1)
808 return {rewriter.getIndexAttr(size), operand};
809 }
810
811 // Filter operands with dynamic dimension
812 auto operandsWithDynamicDim =
813 llvm::filter_to_vector(operands, [&](Value operand) {
814 return cast<RankedTensorType>(operand.getType()).isDynamicDim(dim);
815 });
816
817 // If no operand has a dynamic dimension, it means all sizes were 1
818 if (operandsWithDynamicDim.empty())
819 return {rewriter.getIndexAttr(1), operands.front()};
820
821 // Emit code that computes the runtime size for this dimension. If there is
822 // only one operand with a dynamic dimension, it is considered the master
823 // operand that determines the runtime size of the output dimension.
824 auto targetSize =
825 getTensorDim(rewriter, loc, indexPool, operandsWithDynamicDim[0], dim);
826 if (operandsWithDynamicDim.size() == 1)
827 return {targetSize, operandsWithDynamicDim[0]};
828
829 // Calculate maximum size among all dynamic dimensions
830 for (size_t i = 1; i < operandsWithDynamicDim.size(); i++) {
831 auto nextSize =
832 getTensorDim(rewriter, loc, indexPool, operandsWithDynamicDim[i], dim);
833 targetSize = arith::MaxUIOp::create(rewriter, loc, targetSize, nextSize);
834 }
835 return {targetSize, nullptr};
836}
837
838// Compute the runtime output size for all dimensions. This function returns
839// a pair {targetShape, masterOperands}.
840static std::pair<SmallVector<OpFoldResult>, SmallVector<Value>>
842 IndexPool &indexPool, ValueRange operands) {
843 assert(!operands.empty());
844 auto rank = cast<RankedTensorType>(operands.front().getType()).getRank();
845 SmallVector<OpFoldResult> targetShape;
846 SmallVector<Value> masterOperands;
847 for (auto dim : llvm::seq<int64_t>(0, rank)) {
848 auto [targetSize, masterOperand] =
849 computeTargetSize(rewriter, loc, indexPool, operands, dim);
850 targetShape.push_back(targetSize);
851 masterOperands.push_back(masterOperand);
852 }
853 return {targetShape, masterOperands};
854}
855
857 IndexPool &indexPool, Value operand,
858 int64_t dim, OpFoldResult targetSize,
859 Value masterOperand) {
860 // Nothing to do if this is a static dimension
861 auto rankedTensorType = cast<RankedTensorType>(operand.getType());
862 if (!rankedTensorType.isDynamicDim(dim))
863 return operand;
864
865 // If the target size for this dimension was directly inferred by only taking
866 // this operand into account, there is no need to broadcast. This is an
867 // optimization that will prevent redundant control flow, and constitutes the
868 // main motivation for tracking "master operands".
869 if (operand == masterOperand)
870 return operand;
871
872 // Affine maps for 'linalg.generic' op
873 auto rank = rankedTensorType.getRank();
874 SmallVector<AffineExpr> affineExprs;
875 for (auto index : llvm::seq<int64_t>(0, rank)) {
876 auto affineExpr = index == dim ? rewriter.getAffineConstantExpr(0)
877 : rewriter.getAffineDimExpr(index);
878 affineExprs.push_back(affineExpr);
879 }
880 auto broadcastAffineMap =
881 AffineMap::get(rank, 0, affineExprs, rewriter.getContext());
882 auto identityAffineMap = rewriter.getMultiDimIdentityMap(rank);
883 SmallVector<AffineMap> affineMaps = {broadcastAffineMap, identityAffineMap};
884
885 // Check if broadcast is necessary
886 auto one = createIndex(rewriter, loc, indexPool, 1);
887 auto runtimeSize = getTensorDim(rewriter, loc, indexPool, operand, dim);
888 auto broadcastNecessary = arith::CmpIOp::create(
889 rewriter, loc, arith::CmpIPredicate::eq, runtimeSize, one);
890
891 // Emit 'then' region of 'scf.if'
892 auto emitThenRegion = [&](OpBuilder &opBuilder, Location loc) {
893 // It is not safe to cache constants across regions.
894 // New constants could potentially violate dominance requirements.
895 IndexPool localPool;
896
897 // Emit 'tensor.empty' op
898 SmallVector<OpFoldResult> outputTensorShape;
899 for (auto index : llvm::seq<int64_t>(0, rank)) {
900 auto size = index == dim ? targetSize
901 : getOrFoldTensorDim(rewriter, loc, localPool,
902 operand, index);
903 outputTensorShape.push_back(size);
904 }
905 Value outputTensor = tensor::EmptyOp::create(
906 opBuilder, loc, outputTensorShape, rankedTensorType.getElementType());
907
908 // Emit 'linalg.generic' op
909 auto resultTensor =
910 linalg::GenericOp::create(
911 opBuilder, loc, outputTensor.getType(), operand, outputTensor,
912 affineMaps, getNParallelLoopsAttrs(rank),
913 [&](OpBuilder &opBuilder, Location loc, ValueRange blockArgs) {
914 // Emit 'linalg.yield' op
915 linalg::YieldOp::create(opBuilder, loc, blockArgs.front());
916 })
917 .getResult(0);
918
919 // Cast to original operand type if necessary
920 auto castResultTensor = rewriter.createOrFold<tensor::CastOp>(
921 loc, operand.getType(), resultTensor);
922
923 // Emit 'scf.yield' op
924 scf::YieldOp::create(opBuilder, loc, castResultTensor);
925 };
926
927 // Emit 'else' region of 'scf.if'
928 auto emitElseRegion = [&](OpBuilder &opBuilder, Location loc) {
929 scf::YieldOp::create(opBuilder, loc, operand);
930 };
931
932 // Emit 'scf.if' op
933 auto ifOp = scf::IfOp::create(rewriter, loc, broadcastNecessary,
934 emitThenRegion, emitElseRegion);
935 return ifOp.getResult(0);
936}
937
939 IndexPool &indexPool, Value operand,
940 ArrayRef<OpFoldResult> targetShape,
941 ArrayRef<Value> masterOperands) {
942 int64_t rank = cast<RankedTensorType>(operand.getType()).getRank();
943 assert((int64_t)targetShape.size() == rank);
944 assert((int64_t)masterOperands.size() == rank);
945 for (auto index : llvm::seq<int64_t>(0, rank))
946 operand =
947 broadcastDynamicDimension(rewriter, loc, indexPool, operand, index,
948 targetShape[index], masterOperands[index]);
949 return operand;
950}
951
954 IndexPool &indexPool, ValueRange operands,
955 ArrayRef<OpFoldResult> targetShape,
956 ArrayRef<Value> masterOperands) {
957 // No need to broadcast for unary operations
958 if (operands.size() == 1)
959 return operands;
960
961 // No need to broadcast for static shape
962 bool hasDynamic = false;
963 for (auto op : operands) {
964 const auto tType = dyn_cast<RankedTensorType>(op.getType());
965 if (tType && !tType.hasStaticShape()) {
966 hasDynamic = true;
967 break;
968 }
969 }
970 if (!hasDynamic)
971 return operands;
972
973 // Broadcast dynamic dimensions operand by operand
974 return llvm::map_to_vector(operands, [&](Value operand) {
975 return broadcastDynamicDimensions(rewriter, loc, indexPool, operand,
976 targetShape, masterOperands);
977 });
978}
979
980static LogicalResult
981emitElementwiseComputation(ConversionPatternRewriter &rewriter, Location loc,
982 Operation *operation, ValueRange operands,
983 ArrayRef<OpFoldResult> targetShape,
984 const TypeConverter &converter) {
985 // Generate output tensor
986 auto resultType = cast_or_null<RankedTensorType>(
987 converter.convertType(operation->getResultTypes().front()));
988 if (!resultType) {
989 return rewriter.notifyMatchFailure(operation, "failed to convert type");
990 }
991 Value outputTensor = tensor::EmptyOp::create(rewriter, loc, targetShape,
992 resultType.getElementType());
993
994 // Create affine maps. Input affine maps broadcast static dimensions of size
995 // 1. The output affine map is an identity map.
996 //
997 auto rank = resultType.getRank();
998 auto affineMaps = llvm::map_to_vector(operands, [&](Value operand) {
999 auto shape = cast<ShapedType>(operand.getType()).getShape();
1000 SmallVector<AffineExpr> affineExprs;
1001 for (auto it : llvm::enumerate(shape)) {
1002 // Prefer producting identity maps whenever possible (i.e. no broadcasting
1003 // needed) because some transforms (like reshape folding)
1004 // do not support affine constant exprs.
1005 bool requiresBroadcast =
1006 (it.value() == 1 && resultType.getDimSize(it.index()) != 1);
1007 auto affineExpr = requiresBroadcast
1008 ? rewriter.getAffineConstantExpr(0)
1009 : rewriter.getAffineDimExpr(it.index());
1010 affineExprs.push_back(affineExpr);
1011 }
1012 return AffineMap::get(rank, 0, affineExprs, rewriter.getContext());
1013 });
1014 affineMaps.push_back(rewriter.getMultiDimIdentityMap(rank));
1015
1016 // Emit 'linalg.generic' op
1017 bool encounteredError = false;
1018 auto linalgOp = linalg::GenericOp::create(
1019 rewriter, loc, outputTensor.getType(), operands, outputTensor, affineMaps,
1021 [&](OpBuilder &opBuilder, Location loc, ValueRange blockArgs) {
1023 operation, blockArgs.take_front(operation->getNumOperands()),
1024 {resultType.getElementType()}, rewriter);
1025 if (!opResult) {
1026 encounteredError = true;
1027 return;
1028 }
1029 linalg::YieldOp::create(opBuilder, loc, opResult);
1030 });
1031 if (encounteredError)
1032 return rewriter.notifyMatchFailure(
1033 operation, "unable to create linalg.generic body for elementwise op");
1034
1035 // Cast 'linalg.generic' result into original result type if needed
1036 auto castResult = rewriter.createOrFold<tensor::CastOp>(
1037 loc, resultType, linalgOp->getResult(0));
1038 rewriter.replaceOp(operation, castResult);
1039 return success();
1040}
1041
1043 ValueRange operands) {
1044 // Shift cannot broadcast
1045 if (isa<tosa::MulOp>(operation)) {
1046 DenseElementsAttr shiftElems;
1047 // Shift cannot broadcast when it is constant
1048 if (matchPattern(operation->getOperand(2), m_Constant(&shiftElems)))
1049 return operands.take_front(2);
1050 else
1051 return operands.take_front(3);
1052 }
1053 if (auto negate = dyn_cast<tosa::NegateOp>(operation)) {
1054 FailureOr<int64_t> maybeInZp = negate.getInput1ZeroPoint();
1055 FailureOr<int64_t> maybeOutZp = negate.getOutputZeroPoint();
1056 if (failed(maybeOutZp) && failed(maybeInZp))
1057 return operands;
1058 // Input1_zp and output_zp cannot broadcast when they are constants.
1059 return operands.take_front(1);
1060 }
1061 return operands;
1062}
1063
1064static LogicalResult
1066 ConversionPatternRewriter &rewriter,
1067 const TypeConverter &converter) {
1068
1069 // Collect op properties
1070 assert(operation->getNumResults() == 1 && "elementwise op expects 1 result");
1071 assert(operation->getNumOperands() >= 1 &&
1072 "elementwise op expects at least 1 operand");
1073 if (!operandsAndResultsRanked(operation))
1074 return rewriter.notifyMatchFailure(operation,
1075 "Unranked tensors not supported");
1076
1077 // Lower operation
1078 IndexPool indexPool;
1079 auto loc = operation->getLoc();
1080 auto operandsToBroadcast = getBroadcastableOperands(operation, operands);
1081 auto [targetShape, masterOperands] =
1082 computeTargetShape(rewriter, loc, indexPool, operandsToBroadcast);
1083 auto broadcastOperands =
1084 broadcastDynamicDimensions(rewriter, loc, indexPool, operandsToBroadcast,
1085 targetShape, masterOperands);
1086 return emitElementwiseComputation(rewriter, loc, operation, broadcastOperands,
1087 targetShape, converter);
1088}
1089
1090// Returns the identity value to seed a float min/max reduction with. TOSA seeds
1091// REDUCE_MIN with maximum_s<in_out_t>() and REDUCE_MAX/ARGMAX with
1092// minimum_s<in_out_t>(), and for floating-point types those bounds are
1093// +/-infinity rather than the largest finite value. Only use them when the
1094// caller opted in *and* the format can represent them: APFloat::getInf() is
1095// unreachable for FiniteOnly semantics and silently returns a NaN for NanOnly
1096// semantics such as f8E4M3FN, which would poison the whole reduction through
1097// NaN-propagating arith.minimumf/arith.maximumf.
1098static APFloat getFloatMinMaxIdentity(const llvm::fltSemantics &semantics,
1099 bool negative, bool allowNonFinites) {
1100 if (allowNonFinites && APFloat::semanticsHasInf(semantics))
1101 return APFloat::getInf(semantics, negative);
1102 return APFloat::getLargest(semantics, negative);
1103}
1104
1105// Returns the constant initial value for a given reduction operation. The
1106// attribute type varies depending on the element type required.
1107static TypedAttr createInitialValueForReduceOp(Operation *op, Type elementTy,
1108 PatternRewriter &rewriter,
1109 bool allowNonFinites) {
1110 if (isa<tosa::ReduceSumOp>(op) && isa<FloatType>(elementTy))
1111 return rewriter.getFloatAttr(elementTy, 0.0);
1112
1113 if (isa<tosa::ReduceSumOp>(op) && isa<IntegerType>(elementTy))
1114 return rewriter.getIntegerAttr(elementTy, 0);
1115
1116 if (isa<tosa::ReduceProductOp>(op) && isa<FloatType>(elementTy))
1117 return rewriter.getFloatAttr(elementTy, 1.0);
1118
1119 if (isa<tosa::ReduceProductOp>(op) && isa<IntegerType>(elementTy))
1120 return rewriter.getIntegerAttr(elementTy, 1);
1121
1122 if (isa<tosa::ReduceMinOp>(op) && isa<FloatType>(elementTy))
1123 return rewriter.getFloatAttr(
1124 elementTy,
1125 getFloatMinMaxIdentity(cast<FloatType>(elementTy).getFloatSemantics(),
1126 /*negative=*/false, allowNonFinites));
1127
1128 if (isa<tosa::ReduceMinOp>(op) && isa<IntegerType>(elementTy))
1129 return rewriter.getIntegerAttr(
1130 elementTy, APInt::getSignedMaxValue(elementTy.getIntOrFloatBitWidth()));
1131
1132 if (isa<tosa::ReduceMaxOp>(op) && isa<FloatType>(elementTy))
1133 return rewriter.getFloatAttr(
1134 elementTy,
1135 getFloatMinMaxIdentity(cast<FloatType>(elementTy).getFloatSemantics(),
1136 /*negative=*/true, allowNonFinites));
1137
1138 if (isa<tosa::ReduceMaxOp>(op) && isa<IntegerType>(elementTy))
1139 return rewriter.getIntegerAttr(
1140 elementTy, APInt::getSignedMinValue(elementTy.getIntOrFloatBitWidth()));
1141
1142 if (isa<tosa::ReduceAllOp>(op) && elementTy.isInteger(1))
1143 return rewriter.getIntegerAttr(elementTy, APInt::getAllOnes(1));
1144
1145 if (isa<tosa::ReduceAnyOp>(op) && elementTy.isInteger(1))
1146 return rewriter.getIntegerAttr(elementTy, APInt::getZero(1));
1147
1148 if (isa<tosa::ArgMaxOp>(op) && isa<FloatType>(elementTy))
1149 return rewriter.getFloatAttr(
1150 elementTy,
1151 getFloatMinMaxIdentity(cast<FloatType>(elementTy).getFloatSemantics(),
1152 /*negative=*/true, allowNonFinites));
1153
1154 if (isa<tosa::ArgMaxOp>(op) && isa<IntegerType>(elementTy))
1155 return rewriter.getIntegerAttr(
1156 elementTy, APInt::getSignedMinValue(elementTy.getIntOrFloatBitWidth()));
1157
1158 return {};
1159}
1160
1161// Creates the body calculation for a reduction. The operations vary depending
1162// on the input type.
1164 ValueRange args,
1165 Type elementTy,
1166 PatternRewriter &rewriter) {
1167 Location loc = op->getLoc();
1168 if (isa<tosa::ReduceSumOp>(op) && isa<FloatType>(elementTy)) {
1170 rewriter, loc, TypeRange{elementTy}, args);
1171 }
1172
1173 if (isa<tosa::ReduceSumOp>(op) && isa<IntegerType>(elementTy)) {
1175 rewriter, loc, TypeRange{elementTy}, args);
1176 }
1177
1178 if (isa<tosa::ReduceProductOp>(op) && isa<FloatType>(elementTy)) {
1180 rewriter, loc, TypeRange{elementTy}, args);
1181 }
1182
1183 if (isa<tosa::ReduceProductOp>(op) && isa<IntegerType>(elementTy)) {
1185 rewriter, loc, TypeRange{elementTy}, args);
1186 }
1187
1188 if (isa<tosa::ReduceMinOp>(op) && isa<FloatType>(elementTy)) {
1189 return arith::MinimumFOp::create(rewriter, loc, args[0], args[1]);
1190 }
1191
1192 if (isa<tosa::ReduceMinOp>(op) && isa<IntegerType>(elementTy)) {
1193 return arith::MinSIOp::create(rewriter, loc, args[0], args[1]);
1194 }
1195
1196 if (isa<tosa::ReduceMaxOp>(op) && isa<FloatType>(elementTy)) {
1197 return arith::MaximumFOp::create(rewriter, loc, args[0], args[1]);
1198 }
1199
1200 if (isa<tosa::ReduceMaxOp>(op) && isa<IntegerType>(elementTy)) {
1201 return arith::MaxSIOp::create(rewriter, loc, args[0], args[1]);
1202 }
1203
1204 if (isa<tosa::ReduceAllOp>(op) && elementTy.isInteger(1))
1205 return arith::AndIOp::create(rewriter, loc, args);
1206
1207 if (isa<tosa::ReduceAnyOp>(op) && elementTy.isInteger(1))
1208 return arith::OrIOp::create(rewriter, loc, args);
1209
1210 return {};
1211}
1212
1213// Performs the match and rewrite for reduction operations. This includes
1214// declaring a correctly sized initial value, and the linalg.generic operation
1215// that reduces across the specified axis.
1216template <typename OpTy>
1217static LogicalResult reduceMatchAndRewriteHelper(OpTy op, uint64_t axis,
1218 PatternRewriter &rewriter,
1219 bool allowNonFinites) {
1220 auto loc = op->getLoc();
1221 auto inputTy = dyn_cast<RankedTensorType>(op->getOperand(0).getType());
1222 auto resultTy = dyn_cast<RankedTensorType>(op->getResult(0).getType());
1223 if (!inputTy || !resultTy)
1224 return rewriter.notifyMatchFailure(op, "unranked tensors not supported");
1225
1226 auto elementTy = resultTy.getElementType();
1227 Value input = op->getOperand(0);
1228
1229 // Figure out the accType if needed
1230 bool widenAccTy = std::is_same_v<OpTy, tosa::ReduceSumOp> &&
1231 isa<FloatType>(elementTy) &&
1232 cast<FloatType>(elementTy).isBF16();
1233 Type accTy = widenAccTy ? rewriter.getF32Type() : elementTy;
1234
1235 SmallVector<int64_t> reduceShape;
1236 SmallVector<Value> dynDims;
1237 for (unsigned i = 0; i < inputTy.getRank(); i++) {
1238 if (axis != i) {
1239 reduceShape.push_back(inputTy.getDimSize(i));
1240 if (inputTy.isDynamicDim(i))
1241 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
1242 }
1243 }
1244
1245 SmallVector<Value> inputs, outputs;
1246 inputs.push_back(input);
1247
1248 // First fill the output buffer with the init value.
1249 auto emptyTensor =
1250 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1251 .getResult();
1252
1253 auto fillValueAttr =
1254 createInitialValueForReduceOp(op, accTy, rewriter, allowNonFinites);
1255 if (!fillValueAttr)
1256 return rewriter.notifyMatchFailure(
1257 op, "No initial value found for reduction operation");
1258
1259 auto fillValue = arith::ConstantOp::create(rewriter, loc, fillValueAttr);
1260 auto filledTensor =
1261 linalg::FillOp::create(rewriter, loc, ValueRange{fillValue},
1262 ValueRange{emptyTensor})
1263 .result();
1264 outputs.push_back(filledTensor);
1265
1266 bool isNanIgnoreMode = false;
1267 if constexpr (std::is_same_v<OpTy, tosa::ReduceMinOp> ||
1268 std::is_same_v<OpTy, tosa::ReduceMaxOp>) {
1269 // NaN propagation has no meaning for non floating point types.
1270 if (isa<FloatType>(elementTy) &&
1271 op.getNanMode() == NanPropagationMode::IGNORE) {
1272 isNanIgnoreMode = true;
1273 // Because the TOSA spec requires the result be NaN iff all elements in
1274 // the reduction are NaN we can't simply perform a compare and select.
1275 // Additionally we have to keep track of whether we've seen any non-NaN
1276 // values and then do a final select based on this predicate.
1277 auto trueAttr = rewriter.getBoolAttr(true);
1278 auto trueValue = arith::ConstantOp::create(rewriter, loc, trueAttr);
1279 auto emptyBoolTensor =
1280 tensor::EmptyOp::create(rewriter, loc, reduceShape,
1281 trueValue.getType(), dynDims)
1282 .getResult();
1283 auto allResultsNaNTensor =
1284 linalg::FillOp::create(rewriter, loc, ValueRange{trueValue},
1285 ValueRange{emptyBoolTensor})
1286 .result();
1287 // Note that because the linalg::ReduceOp has two variadic arguments
1288 // (inputs and outputs) and it has the SameVariadicOperandSize trait we
1289 // need to have the same number of inputs and outputs.
1290 //
1291 // The second input isn't actually used anywhere since the value used to
1292 // update the NaN flag is calculated inside the body of the reduction and
1293 // then used to update an out value.
1294 // In order to satisfy type constraints we just pass another copy of the
1295 // input here.
1296 inputs.push_back(input);
1297 outputs.push_back(allResultsNaNTensor);
1298 }
1299 }
1300
1301 bool didEncounterError = false;
1302 linalg::LinalgOp linalgOp = linalg::ReduceOp::create(
1303 rewriter, loc, inputs, outputs, axis,
1304 [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange blockArgs) {
1305 std::array<Value, 2> binaryArgs{
1306 blockArgs[0], isNanIgnoreMode ? blockArgs[2] : blockArgs[1]};
1307
1308 // If reduction type differs then extend (applicable to reduce_sum)
1309 if (binaryArgs[0].getType() != accTy)
1310 binaryArgs[0] = arith::ExtFOp::create(
1311 nestedBuilder, nestedLoc, TypeRange{accTy},
1312 ValueRange{binaryArgs[0]}, arith::ExtFOp::Properties{});
1313
1314 auto result = createLinalgBodyCalculationForReduceOp(op, binaryArgs,
1315 accTy, rewriter);
1316 if (result)
1317 didEncounterError = true;
1318
1319 SmallVector<Value> resultsToYield;
1320 if (isNanIgnoreMode) {
1321 auto inputValue = blockArgs[0];
1322 auto initialValue = blockArgs[2];
1323 auto oldAllResultsNanFlagValue = blockArgs[3];
1324
1325 // Unordered comparison of NaN against itself will always return true.
1326 Value isNaN = arith::CmpFOp::create(nestedBuilder, op->getLoc(),
1327 arith::CmpFPredicate::UNO,
1328 inputValue, inputValue);
1329 // If we've encountered a NaN, take the non-NaN value.
1330 auto selectOp = arith::SelectOp::create(nestedBuilder, op->getLoc(),
1331 isNaN, initialValue, result);
1332 // Update the flag which keeps track of whether we have seen a non-NaN
1333 // value.
1334 auto newAllResultsNanFlagValue = arith::AndIOp::create(
1335 nestedBuilder, op->getLoc(), oldAllResultsNanFlagValue, isNaN);
1336 resultsToYield.push_back(selectOp);
1337 resultsToYield.push_back(newAllResultsNanFlagValue);
1338 } else {
1339 resultsToYield.push_back(result);
1340 }
1341 linalg::YieldOp::create(nestedBuilder, loc, resultsToYield);
1342 });
1343
1344 if (!didEncounterError)
1345 return rewriter.notifyMatchFailure(
1346 op, "unable to create linalg.generic body for reduce op");
1347
1348 if (isNanIgnoreMode) {
1349 // Materialize a check to see whether we encountered any non-NaN values, if
1350 // we didn't we need to select a tensor of NaNs since the result will just
1351 // be the initial identity value propagated through all the compares and
1352 // selects inside the reduction.
1353
1354 // Create a tensor full of NaNs.
1355 auto nanValueAttr = rewriter.getFloatAttr(
1356 accTy,
1357 APFloat::getNaN(cast<FloatType>(elementTy).getFloatSemantics(), false));
1358 auto nanValue = arith::ConstantOp::create(rewriter, loc, nanValueAttr);
1359 auto emptyNanTensor =
1360 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1361 .getResult();
1362 auto nanFilledTensor =
1363 linalg::FillOp::create(rewriter, loc, ValueRange{nanValue},
1364 ValueRange{emptyNanTensor})
1365 .result();
1366
1367 // Create an empty tensor, non need to fill this since it will be
1368 // overwritten by the select.
1369 auto finalEmptyTensor =
1370 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1371 .getResult();
1372
1373 // Do a selection between the tensors akin to:
1374 // result = NaN if "all results NaN" else result.
1375 SmallVector<Value> ins, outs;
1376 ins.push_back(linalgOp->getOpResult(1));
1377 ins.push_back(nanFilledTensor);
1378 ins.push_back(linalgOp->getResult(0));
1379 outs.push_back(finalEmptyTensor);
1380 auto linalgSelect =
1381 linalg::ElementwiseOp::create(rewriter, op->getLoc(), ins, outs,
1382 mlir::linalg::ElementwiseKind::select);
1383 linalgOp = linalgSelect;
1384 }
1385
1386 // Truncate back to resultTy if needed
1387 Value reducedRes = linalgOp->getResult(0);
1388 if (widenAccTy) {
1389 auto resEmptyOp =
1390 tensor::EmptyOp::create(rewriter, loc, reduceShape, elementTy, dynDims)
1391 .getResult();
1392
1393 const unsigned reducedRank =
1394 cast<ShapedType>(reducedRes.getType()).getRank();
1395 auto identityMap = rewriter.getMultiDimIdentityMap(reducedRank);
1396 reducedRes =
1397 linalg::GenericOp::create(
1398 rewriter, loc, resEmptyOp.getType(), ValueRange{reducedRes},
1399 ValueRange{resEmptyOp},
1400 ArrayRef<AffineMap>{identityMap, identityMap},
1401 getNParallelLoopsAttrs(reducedRank),
1402 [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange args) {
1403 Value truncf = arith::TruncFOp::create(nestedBuilder, nestedLoc,
1404 elementTy, args[0]);
1405 linalg::YieldOp::create(nestedBuilder, nestedLoc, truncf);
1406 })
1407 .getResults()[0];
1408 }
1409
1410 SmallVector<ReassociationExprs, 4> reassociationMap;
1411 uint64_t expandInputRank = cast<ShapedType>(reducedRes.getType()).getRank();
1412 reassociationMap.resize(expandInputRank);
1413
1414 for (uint64_t i = 0; i < expandInputRank; i++) {
1415 int32_t dimToPush = i > axis ? i + 1 : i;
1416 reassociationMap[i].push_back(rewriter.getAffineDimExpr(dimToPush));
1417 }
1418
1419 if (expandInputRank != 0) {
1420 int32_t expandedDim = axis < expandInputRank ? axis : expandInputRank - 1;
1421 reassociationMap[expandedDim].push_back(
1422 rewriter.getAffineDimExpr(expandedDim + 1));
1423 }
1424
1425 // Lower directly to `tensor::ExpandShapeOp` instead of `tosa::ReshapeOp`,
1426 // since here we know which dimension to expand, and `tosa::ReshapeOp` would
1427 // not have access to such information. This matters when handling dynamically
1428 // sized tensors.
1429 rewriter.replaceOpWithNewOp<tensor::ExpandShapeOp>(op, resultTy, reducedRes,
1430 reassociationMap);
1431 return success();
1432}
1433
1434namespace {
1435
1436template <typename SrcOp>
1437class PointwiseConverter : public OpConversionPattern<SrcOp> {
1438public:
1439 using OpConversionPattern<SrcOp>::OpConversionPattern;
1440 using typename OpConversionPattern<SrcOp>::OpAdaptor;
1441
1442 LogicalResult
1443 matchAndRewrite(SrcOp op, OpAdaptor operands,
1444 ConversionPatternRewriter &rewriter) const final {
1446 op, operands.getOperands(), rewriter, *this->getTypeConverter());
1447 }
1448};
1449
1450// Collapse tensor<1xiN> into tensor<iN>
1451// E.g. tensor.collapse_shape %arg1 [] : tensor<1xi16> into tensor<i16>
1452static Value collapse1xNTensorToN(PatternRewriter &rewriter, Value input,
1453 Location loc) {
1455 // Create the collapsed type
1456 auto inputType = cast<RankedTensorType>(input.getType());
1457 auto elemType = inputType.getElementType();
1458 auto collapsedType = RankedTensorType::get({}, elemType);
1459 // Emit the collapse op
1460 return tensor::CollapseShapeOp::create(rewriter, loc, collapsedType, input,
1461 reassociation);
1462}
1463
1465convertToI8(const llvm::SmallVector<int32_t> &input) {
1467 output.reserve(input.size());
1468
1469 for (auto v : llvm::map_range(
1470 input, [](int32_t val) { return static_cast<int8_t>(val); })) {
1471 output.push_back(v);
1472 }
1473 return output;
1474}
1475
1476// The shift or multiplier may be either constant or non-constant, depending on
1477// whether dynamic extension is enabled.
1478// - If the shift or multiplier is non-constant, add it as an input to
1479// linalg::GenericOp by:
1480// 1. Pushing it into 'genericInputs'.
1481// 2. Appending a corresponding affine map to 'indexingMaps'.
1482// - If the shift or multiplier is constant, set 'constant' instead.
1483static void setupLinalgGenericOpInputAndIndexingMap(
1485 SmallVector<Value, 4> &genericInputs, SmallVector<AffineMap> &indexingMaps,
1486 bool isConstant, tosa::RescaleOp op, Value &constant, int64_t &arg,
1487 bool isShift = false) {
1488
1489 auto loc = op.getLoc();
1490 auto inputTy = cast<ShapedType>(op.getInput().getType());
1491 unsigned rank = inputTy.getRank();
1492 SmallVector<AffineExpr, 2> exprs = {rewriter.getAffineDimExpr(rank - 1)};
1493
1494 if (isConstant) {
1495 // If we are rescaling per-channel then we need to store the
1496 // values in a buffer.
1497 if (values.size() == 1) {
1498 IntegerAttr intAttr = isShift
1499 ? rewriter.getI8IntegerAttr(values.front())
1500 : rewriter.getI32IntegerAttr(values.front());
1501 constant = arith::ConstantOp::create(rewriter, loc, intAttr);
1502 } else {
1503 auto elementType =
1504 isShift ? rewriter.getIntegerType(8) : rewriter.getI32Type();
1505 auto tensorType = RankedTensorType::get(
1506 {static_cast<int64_t>(values.size())}, elementType);
1507 DenseIntElementsAttr EltAttr;
1508 if (isShift)
1509 EltAttr = DenseIntElementsAttr::get(tensorType, convertToI8(values));
1510 else
1511 EltAttr = DenseIntElementsAttr::get(tensorType, values);
1512 genericInputs.push_back(
1513 arith::ConstantOp::create(rewriter, loc, EltAttr));
1514 indexingMaps.push_back(AffineMap::get(/*dimCount=*/rank,
1515 /*symbolCount=*/0, exprs,
1516 rewriter.getContext()));
1517 }
1518 } else {
1519 // If we are not rescaling per-channel then we need to collapse 1xN to N
1520 // and push broadcastMap.
1521 auto operand = isShift ? op.getShift() : op.getMultiplier();
1522 auto tensorType = dyn_cast<RankedTensorType>(operand.getType());
1523 if (tensorType && tensorType.hasStaticShape() &&
1524 tensorType.getShape()[0] == 1) {
1525 // broadcastMap = affine_map<(d0, d1) -> ()>
1526 // It would affect as broadcast for scalar values in linalg::GenericOp.
1527 AffineMap broadcastMap =
1528 AffineMap::get(rank, 0, {}, rewriter.getContext());
1529 genericInputs.push_back(collapse1xNTensorToN(rewriter, operand, loc));
1530 indexingMaps.push_back(broadcastMap);
1531 } else {
1532 genericInputs.push_back(operand);
1533 indexingMaps.push_back(AffineMap::get(/*dimCount=*/rank,
1534 /*symbolCount=*/0, exprs,
1535 rewriter.getContext()));
1536 }
1537 }
1538 arg = indexingMaps.size() - 1;
1539}
1540
1541// Return the extended Zp to be used in subsequent arithmetic operations.
1542static Value getExtendZp(OpBuilder &builder, Type valueTy,
1543 FailureOr<int64_t> maybeZp, Location loc,
1544 ValueRange blockArgs, int64_t zpArg,
1545 bool isOutputZp = false) {
1546 Value result;
1547 const int32_t bitwidth = valueTy.getIntOrFloatBitWidth();
1548 const uint32_t attrBitwidth =
1549 isOutputZp ? 32 : (bitwidth > 32 ? bitwidth : 32);
1550 auto extendType = builder.getIntegerType(attrBitwidth);
1551 // The Zp value can be either constant or non-constant, depending on
1552 // whether dynamic extension is enabled.
1553 // If 'maybeZp' fails, it indicates that Zp is non-constant and will
1554 // be passed as an input to linalg::GenericOp.
1555 if (failed(maybeZp)) {
1556 result = blockArgs[zpArg];
1557 auto zpTy = result.getType();
1558 if (zpTy.getIntOrFloatBitWidth() < attrBitwidth) {
1559 // For ExtUIOp, the input must be signless.
1560 // UnrealizedConversionCastOp will cast the input to signless type.
1561 if (zpTy.isUnsignedInteger()) {
1562 result =
1563 UnrealizedConversionCastOp::create(
1564 builder, loc,
1565 builder.getIntegerType(zpTy.getIntOrFloatBitWidth()), result)
1566 .getResult(0);
1567 }
1568 if (zpTy.isUnsignedInteger()) {
1569 return arith::ExtUIOp::create(builder, loc, extendType, result);
1570 } else {
1571 return arith::ExtSIOp::create(builder, loc, extendType, result);
1572 }
1573 }
1574 } else {
1575 return arith::ConstantOp::create(builder, loc,
1576 IntegerAttr::get(extendType, *maybeZp));
1577 }
1578 return result;
1579}
1580
1581class RescaleConverter : public OpRewritePattern<tosa::RescaleOp> {
1582public:
1583 using OpRewritePattern<tosa::RescaleOp>::OpRewritePattern;
1584
1585 LogicalResult matchAndRewrite(tosa::RescaleOp op,
1586 PatternRewriter &rewriter) const final {
1587 auto loc = op.getLoc();
1588 auto input = op.getInput();
1589 auto inputTy = cast<ShapedType>(op.getInput().getType());
1590 auto outputTy = cast<ShapedType>(op.getOutput().getType());
1591 unsigned rank = inputTy.getRank();
1592
1593 // This is an illegal configuration. terminate and log an error
1594 if (op.getRoundingMode() == RoundingMode::INEXACT_ROUND)
1595 return rewriter.notifyMatchFailure(
1596 op, "tosa.rescale with rounding mode = 'INEXACT_ROUND' is not "
1597 "currently supported");
1598 if (op.getRoundingMode() == RoundingMode::DOUBLE_ROUND && !op.getScale32())
1599 return rewriter.notifyMatchFailure(
1600 op, "tosa.rescale requires scale32 for double_round to be true");
1601
1602 if (!isa<IntegerType>(inputTy.getElementType()))
1603 return rewriter.notifyMatchFailure(op, "only support integer type");
1604
1605 SmallVector<Value> dynDims;
1606 for (int i = 0; i < outputTy.getRank(); i++) {
1607 if (outputTy.isDynamicDim(i)) {
1608 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
1609 }
1610 }
1611
1612 DenseElementsAttr shiftElems;
1613 bool isShiftConstant = false;
1614 if (matchPattern(op.getShift(), m_Constant(&shiftElems)))
1615 isShiftConstant = true;
1616
1617 DenseElementsAttr multiplierElems;
1618 bool isMultiplierConstant = false;
1619 if (matchPattern(op.getMultiplier(), m_Constant(&multiplierElems)))
1620 isMultiplierConstant = true;
1621
1622 llvm::SmallVector<int32_t> shiftValues;
1623 llvm::SmallVector<int32_t> multiplierValues;
1624 bool doubleRound;
1625
1626 if (isMultiplierConstant && isShiftConstant) {
1627 // explicit cast is required here
1628 shiftValues = llvm::map_to_vector(
1629 shiftElems.getValues<IntegerAttr>(), [](IntegerAttr attr) -> int32_t {
1630 return static_cast<int32_t>(attr.getInt());
1631 });
1632 multiplierValues =
1633 llvm::map_to_vector(multiplierElems.getValues<IntegerAttr>(),
1634 [](IntegerAttr attr) -> int32_t {
1635 return static_cast<int32_t>(attr.getInt());
1636 });
1637
1638 // If we shift by more than the bitwidth, this just sets to 0.
1639 for (int i = 0, s = multiplierValues.size(); i < s; i++) {
1640 if (shiftValues[i] > 63) {
1641 shiftValues[i] = 0;
1642 multiplierValues[i] = 0;
1643 }
1644 }
1645 // Double round only occurs if shift is greater than 31, check that this
1646 // is ever true.
1647 doubleRound = op.getRoundingMode() == RoundingMode::DOUBLE_ROUND &&
1648 llvm::any_of(shiftValues, [](int32_t v) { return v > 31; });
1649 } else
1650 doubleRound = op.getRoundingMode() == RoundingMode::DOUBLE_ROUND;
1651
1652 RoundingMode roundingMode =
1653 doubleRound ? RoundingMode::DOUBLE_ROUND : RoundingMode::SINGLE_ROUND;
1654
1655 SmallVector<AffineMap> indexingMaps = {
1656 rewriter.getMultiDimIdentityMap(rank)};
1657 SmallVector<Value, 4> genericInputs = {input};
1658
1659 // If we are rescaling per-channel then we need to store the multiplier
1660 // values in a buffer.
1661 Value multiplierConstant;
1662 int64_t multiplierArg = 0;
1663 setupLinalgGenericOpInputAndIndexingMap(
1664 rewriter, multiplierValues, genericInputs, indexingMaps,
1665 isMultiplierConstant, op, multiplierConstant, multiplierArg);
1666
1667 // If we are rescaling per-channel then we need to store the shift
1668 // values in a buffer.
1669 Value shiftConstant;
1670 int64_t shiftArg = 0;
1671 setupLinalgGenericOpInputAndIndexingMap(
1672 rewriter, shiftValues, genericInputs, indexingMaps, isShiftConstant, op,
1673 shiftConstant, shiftArg, true);
1674
1675 // broadcastMap = affine_map<(d0, d1) -> ()>
1676 // It would affect as broadcast for scalar values in linalg::GenericOp.
1677 AffineMap broadcastMap = AffineMap::get(rank, 0, {}, rewriter.getContext());
1678 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1679 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1680 // The inputZp and outputZp may be either constant or non-constant,
1681 // depending on whether dynamic extension is enabled.
1682 // - If the zp's are non-constant, add them as an inputs to
1683 // linalg::GenericOp by:
1684 // 1. Pushing it into 'genericInputs'.
1685 // 2. Appending a corresponding affine map to 'indexingMaps'.
1686 // - If the zp's are constant, they would be generated as arith.constant.
1687 int64_t iZpArg = 0;
1688 if (failed(maybeIZp)) {
1689 genericInputs.push_back(
1690 collapse1xNTensorToN(rewriter, op->getOperand(3), loc));
1691 indexingMaps.push_back(broadcastMap);
1692 iZpArg = indexingMaps.size() - 1;
1693 }
1694 int64_t oZpArg = 0;
1695 if (failed(maybeOZp)) {
1696 genericInputs.push_back(
1697 collapse1xNTensorToN(rewriter, op->getOperand(4), loc));
1698 indexingMaps.push_back(broadcastMap);
1699 oZpArg = indexingMaps.size() - 1;
1700 }
1701
1702 // Indexing maps for output values.
1703 indexingMaps.push_back(rewriter.getMultiDimIdentityMap(rank));
1704
1705 // Construct the indexing maps needed for linalg.generic ops.
1706 Value emptyTensor = tensor::EmptyOp::create(
1707 rewriter, loc, outputTy.getShape(), outputTy.getElementType(),
1708 ArrayRef<Value>({dynDims}));
1709
1710 auto linalgOp = linalg::GenericOp::create(
1711 rewriter, loc, outputTy, genericInputs, ValueRange{emptyTensor},
1712 indexingMaps, getNParallelLoopsAttrs(rank),
1713 [&](OpBuilder &nestedBuilder, Location nestedLoc,
1714 ValueRange blockArgs) {
1715 Value value = blockArgs[0];
1716 Type valueTy = value.getType();
1717
1718 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1719 auto inputZp = getExtendZp(nestedBuilder, valueTy, maybeIZp,
1720 nestedLoc, blockArgs, iZpArg);
1721
1722 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1723 auto outputZp = getExtendZp(nestedBuilder, valueTy, maybeOZp,
1724 nestedLoc, blockArgs, oZpArg, true);
1725
1726 IntegerType outIntType =
1727 cast<IntegerType>(blockArgs.back().getType());
1728 unsigned outBitWidth = outIntType.getWidth();
1729 assert(outBitWidth <= 32 && "Unexpected output zeropoint bitwidth");
1730
1731 Value multiplier = multiplierConstant ? multiplierConstant
1732 : blockArgs[multiplierArg];
1733 Value shift = shiftConstant ? shiftConstant : blockArgs[shiftArg];
1734
1735 if (valueTy.isUnsignedInteger()) {
1736 value = UnrealizedConversionCastOp::create(
1737 nestedBuilder, nestedLoc,
1738 nestedBuilder.getIntegerType(
1739 valueTy.getIntOrFloatBitWidth()),
1740 value)
1741 .getResult(0);
1742 }
1743 if (valueTy.getIntOrFloatBitWidth() < 32) {
1744 if (op.getInputUnsigned()) {
1745 value = arith::ExtUIOp::create(nestedBuilder, nestedLoc,
1746 nestedBuilder.getI32Type(), value);
1747 } else {
1748 value = arith::ExtSIOp::create(nestedBuilder, nestedLoc,
1749 nestedBuilder.getI32Type(), value);
1750 }
1751 }
1752
1753 value =
1754 arith::SubIOp::create(nestedBuilder, nestedLoc, value, inputZp);
1755
1756 value = tosa::ApplyScaleOp::create(nestedBuilder, loc,
1757 nestedBuilder.getI32Type(), value,
1758 multiplier, shift, roundingMode);
1759
1760 // Move to the new zero-point.
1761 value =
1762 arith::AddIOp::create(nestedBuilder, nestedLoc, value, outputZp);
1763
1764 // Saturate to the output size.
1765 int32_t intMin = APInt::getSignedMinValue(outBitWidth).getSExtValue();
1766 int32_t intMax = APInt::getSignedMaxValue(outBitWidth).getSExtValue();
1767
1768 // Unsigned integers have a difference output value.
1769 if (op.getOutputUnsigned()) {
1770 intMin = 0;
1771 intMax = APInt::getMaxValue(outBitWidth).getZExtValue();
1772 }
1773
1774 auto intMinVal = arith::ConstantOp::create(
1775 nestedBuilder, loc, nestedBuilder.getI32IntegerAttr(intMin));
1776 auto intMaxVal = arith::ConstantOp::create(
1777 nestedBuilder, loc, nestedBuilder.getI32IntegerAttr(intMax));
1778
1779 value = clampIntHelper(nestedLoc, value, intMinVal, intMaxVal,
1780 nestedBuilder, /*isUnsigned=*/false);
1781
1782 if (outIntType.getWidth() < 32) {
1783 value = arith::TruncIOp::create(
1784 nestedBuilder, nestedLoc,
1785 rewriter.getIntegerType(outIntType.getWidth()), value);
1786 }
1787
1788 if (outIntType.isUnsignedInteger()) {
1789 value = UnrealizedConversionCastOp::create(nestedBuilder, nestedLoc,
1790 outIntType, value)
1791 .getResult(0);
1792 }
1793 linalg::YieldOp::create(nestedBuilder, loc, value);
1794 });
1795
1796 rewriter.replaceOp(op, linalgOp->getResults());
1797 return success();
1798 }
1799};
1800
1801// Handle the resize case where the input is a 1x1 image. This case
1802// can entirely avoiding having extract operations which target much
1803// more difficult to optimize away.
1804class ResizeUnaryConverter : public OpRewritePattern<tosa::ResizeOp> {
1805public:
1806 using OpRewritePattern<tosa::ResizeOp>::OpRewritePattern;
1807
1808 LogicalResult matchAndRewrite(tosa::ResizeOp op,
1809 PatternRewriter &rewriter) const final {
1810 Location loc = op.getLoc();
1811 ImplicitLocOpBuilder builder(loc, rewriter);
1812 auto input = op.getInput();
1813 auto inputTy = cast<RankedTensorType>(input.getType());
1814 auto resultTy = cast<RankedTensorType>(op.getType());
1815 const bool isBilinear = op.getMode() == ResizeMode::BILINEAR;
1816
1817 auto inputH = inputTy.getDimSize(1);
1818 auto inputW = inputTy.getDimSize(2);
1819 auto outputH = resultTy.getDimSize(1);
1820 auto outputW = resultTy.getDimSize(2);
1821
1822 if (inputH != 1 || inputW != 1 || outputH != 1 || outputW != 1)
1823 return rewriter.notifyMatchFailure(
1824 op, "tosa.resize is not a pure 1x1->1x1 image operation");
1825
1826 if (op.getMode() != ResizeMode::NEAREST_NEIGHBOR &&
1827 op.getMode() != ResizeMode::BILINEAR)
1828 return rewriter.notifyMatchFailure(
1829 op, "tosa.resize mode should be NEAREST_NEIGHBOR or BILINEAR");
1830
1831 if (inputTy == resultTy) {
1832 rewriter.replaceOp(op, input);
1833 return success();
1834 }
1835
1836 SmallVector<int64_t> scale;
1837 if (!tosa::getConstShapeValues(op.getScale().getDefiningOp(), scale)) {
1838 return failure();
1839 }
1840
1841 // Collapse the unit width and height away.
1842 SmallVector<ReassociationExprs, 4> reassociationMap(2);
1843 reassociationMap[0].push_back(builder.getAffineDimExpr(0));
1844 reassociationMap[1].push_back(builder.getAffineDimExpr(1));
1845 reassociationMap[1].push_back(builder.getAffineDimExpr(2));
1846 reassociationMap[1].push_back(builder.getAffineDimExpr(3));
1847
1848 auto collapseTy =
1849 RankedTensorType::get({inputTy.getDimSize(0), inputTy.getDimSize(3)},
1850 inputTy.getElementType());
1851 Value collapse = tensor::CollapseShapeOp::create(builder, collapseTy, input,
1852 reassociationMap);
1853
1854 // Get any dynamic shapes that appear in the input format.
1855 llvm::SmallVector<Value> outputDynSize;
1856 if (inputTy.isDynamicDim(0))
1857 outputDynSize.push_back(tensor::DimOp::create(builder, input, 0));
1858 if (inputTy.isDynamicDim(3))
1859 outputDynSize.push_back(tensor::DimOp::create(builder, input, 3));
1860
1861 // Generate the elementwise operation for casting scaling the input value.
1862 auto genericTy = collapseTy.clone(resultTy.getElementType());
1863 Value empty =
1864 tensor::EmptyOp::create(builder, genericTy.getShape(),
1865 resultTy.getElementType(), outputDynSize);
1866 auto genericMap = rewriter.getMultiDimIdentityMap(genericTy.getRank());
1867 SmallVector<utils::IteratorType> iterators(genericTy.getRank(),
1868 utils::IteratorType::parallel);
1869
1870 auto generic = linalg::GenericOp::create(
1871 builder, genericTy, ValueRange{collapse}, ValueRange{empty},
1872 ArrayRef<AffineMap>{genericMap, genericMap}, iterators,
1873 [=](OpBuilder &b, Location loc, ValueRange args) {
1874 Value value = args[0];
1875 // This is the quantized case.
1876 if (inputTy.getElementType() != resultTy.getElementType()) {
1877 value = arith::ExtSIOp::create(b, loc, resultTy.getElementType(),
1878 value);
1879
1880 if (isBilinear && scale[0] != 0) {
1881 Value scaleY = arith::ConstantOp::create(
1882 b, loc, b.getI32IntegerAttr(scale[0]));
1883 value = arith::MulIOp::create(b, loc, value, scaleY);
1884 }
1885
1886 if (isBilinear && scale[2] != 0) {
1887 Value scaleX = arith::ConstantOp::create(
1888 b, loc, b.getI32IntegerAttr(scale[2]));
1889 value = arith::MulIOp::create(b, loc, value, scaleX);
1890 }
1891 }
1892
1893 linalg::YieldOp::create(b, loc, value);
1894 });
1895
1896 rewriter.replaceOpWithNewOp<tensor::ExpandShapeOp>(
1897 op, resultTy, generic.getResults()[0], reassociationMap);
1898 return success();
1899 }
1900};
1901
1902// TOSA resize with width or height of 1 may be broadcasted to a wider
1903// dimension. This is done by materializing a new tosa.resize without
1904// the broadcasting behavior, and an explicit broadcast afterwards.
1905class MaterializeResizeBroadcast : public OpRewritePattern<tosa::ResizeOp> {
1906public:
1907 using OpRewritePattern<tosa::ResizeOp>::OpRewritePattern;
1908
1909 LogicalResult matchAndRewrite(tosa::ResizeOp op,
1910 PatternRewriter &rewriter) const final {
1911 Location loc = op.getLoc();
1912 ImplicitLocOpBuilder builder(loc, rewriter);
1913 auto input = op.getInput();
1914 auto inputTy = dyn_cast<RankedTensorType>(input.getType());
1915 auto resultTy = dyn_cast<RankedTensorType>(op.getType());
1916
1917 if (!inputTy || !resultTy)
1918 return rewriter.notifyMatchFailure(op,
1919 "requires ranked input/output types");
1920
1921 auto batch = inputTy.getDimSize(0);
1922 auto channels = inputTy.getDimSize(3);
1923 auto inputH = inputTy.getDimSize(1);
1924 auto inputW = inputTy.getDimSize(2);
1925 auto outputH = resultTy.getDimSize(1);
1926 auto outputW = resultTy.getDimSize(2);
1927
1928 if ((inputH != 1 || outputH == 1) && (inputW != 1 || outputW == 1))
1929 return rewriter.notifyMatchFailure(
1930 op, "tosa.resize has no broadcasting behavior");
1931
1932 // For any dimension that is broadcastable we generate a width of 1
1933 // on the output.
1934 llvm::SmallVector<int64_t> resizeShape;
1935 resizeShape.push_back(batch);
1936 resizeShape.push_back(inputH == 1 ? 1 : outputH);
1937 resizeShape.push_back(inputW == 1 ? 1 : outputW);
1938 resizeShape.push_back(channels);
1939
1940 auto resizeTy = resultTy.clone(resizeShape);
1941 auto resize =
1942 tosa::ResizeOp::create(builder, resizeTy, input, op.getScale(),
1943 op.getOffset(), op.getBorder(), op.getMode());
1944
1945 // Collapse an unit result dims.
1946 SmallVector<ReassociationExprs, 4> reassociationMap(2);
1947 reassociationMap[0].push_back(builder.getAffineDimExpr(0));
1948 reassociationMap.back().push_back(builder.getAffineDimExpr(1));
1949 if (inputH != 1)
1950 reassociationMap.push_back({});
1951 reassociationMap.back().push_back(builder.getAffineDimExpr(2));
1952 if (inputW != 1)
1953 reassociationMap.push_back({});
1954 reassociationMap.back().push_back(builder.getAffineDimExpr(3));
1955
1956 llvm::SmallVector<int64_t> collapseShape = {batch};
1957 if (inputH != 1)
1958 collapseShape.push_back(outputH);
1959 if (inputW != 1)
1960 collapseShape.push_back(outputW);
1961 collapseShape.push_back(channels);
1962
1963 auto collapseTy = resultTy.clone(collapseShape);
1964 Value collapse = tensor::CollapseShapeOp::create(builder, collapseTy,
1965 resize, reassociationMap);
1966
1967 // Broadcast the collapsed shape to the output result.
1968 llvm::SmallVector<Value> outputDynSize;
1969 if (inputTy.isDynamicDim(0))
1970 outputDynSize.push_back(tensor::DimOp::create(builder, input, 0));
1971 if (inputTy.isDynamicDim(3))
1972 outputDynSize.push_back(tensor::DimOp::create(builder, input, 3));
1973
1974 SmallVector<utils::IteratorType> iterators(resultTy.getRank(),
1975 utils::IteratorType::parallel);
1976 Value empty = tensor::EmptyOp::create(
1977 builder, resultTy.getShape(), resultTy.getElementType(), outputDynSize);
1978
1979 SmallVector<AffineExpr, 4> inputExprs{rewriter.getAffineDimExpr(0)};
1980 if (inputH != 1)
1981 inputExprs.push_back(rewriter.getAffineDimExpr(1));
1982 if (inputW != 1)
1983 inputExprs.push_back(rewriter.getAffineDimExpr(2));
1984 inputExprs.push_back(rewriter.getAffineDimExpr(3));
1985
1986 auto inputMap = AffineMap::get(resultTy.getRank(), /*symbolCount=*/0,
1987 inputExprs, rewriter.getContext());
1988
1989 auto outputMap = rewriter.getMultiDimIdentityMap(resultTy.getRank());
1990 rewriter.replaceOpWithNewOp<linalg::GenericOp>(
1991 op, resultTy, ValueRange{collapse}, ValueRange{empty},
1992 ArrayRef<AffineMap>{inputMap, outputMap}, iterators,
1993 [=](OpBuilder &b, Location loc, ValueRange args) {
1994 Value value = args[0];
1995 linalg::YieldOp::create(b, loc, value);
1996 });
1997
1998 return success();
1999 }
2000};
2001
2002class GenericResizeConverter : public OpRewritePattern<tosa::ResizeOp> {
2003public:
2004 using OpRewritePattern<tosa::ResizeOp>::OpRewritePattern;
2005
2006 LogicalResult matchAndRewrite(tosa::ResizeOp op,
2007 PatternRewriter &rewriter) const final {
2008 Location loc = op.getLoc();
2009 ImplicitLocOpBuilder b(loc, rewriter);
2010 auto input = op.getInput();
2011 auto inputTy = cast<ShapedType>(input.getType());
2012 auto resultTy = cast<ShapedType>(op.getType());
2013 auto resultETy = resultTy.getElementType();
2014
2015 bool floatingPointMode = isa<FloatType>(resultETy);
2016 auto floatTy = resultETy;
2017
2018 auto imageH = inputTy.getShape()[1];
2019 auto imageW = inputTy.getShape()[2];
2020
2021 auto dynamicDimsOr =
2022 checkHasDynamicBatchDims(rewriter, op, {input, op.getOutput()});
2023 if (!dynamicDimsOr.has_value())
2024 return rewriter.notifyMatchFailure(
2025 op, "unable to get dynamic dimensions of tosa.resize");
2026
2027 if (op.getMode() != ResizeMode::NEAREST_NEIGHBOR &&
2028 op.getMode() != ResizeMode::BILINEAR)
2029 return rewriter.notifyMatchFailure(
2030 op, "tosa.resize mode should be NEAREST_NEIGHBOR or BILINEAR");
2031
2032 SmallVector<AffineMap, 2> affineMaps = {
2033 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2034 auto emptyTensor = tensor::EmptyOp::create(b, resultTy.getShape(),
2035 resultETy, *dynamicDimsOr);
2036 auto genericOp = linalg::GenericOp::create(
2037 b, resultTy, ValueRange({}), ValueRange{emptyTensor}, affineMaps,
2038 getNParallelLoopsAttrs(resultTy.getRank()));
2039 Value resize = genericOp.getResult(0);
2040
2041 {
2042 OpBuilder::InsertionGuard regionGuard(b);
2043 b.createBlock(&genericOp.getRegion(), genericOp.getRegion().end(),
2044 TypeRange({resultETy}), loc);
2045 Value batch = linalg::IndexOp::create(b, 0);
2046 Value y = linalg::IndexOp::create(b, 1);
2047 Value x = linalg::IndexOp::create(b, 2);
2048 Value channel = linalg::IndexOp::create(b, 3);
2049
2050 Value zeroI32 =
2051 arith::ConstantOp::create(b, b.getZeroAttr(b.getI32Type()));
2052 Value zeroFp = arith::ConstantOp::create(b, b.getZeroAttr(floatTy));
2053 Value hMax =
2054 arith::ConstantOp::create(b, b.getI32IntegerAttr(imageH - 1));
2055 Value wMax =
2056 arith::ConstantOp::create(b, b.getI32IntegerAttr(imageW - 1));
2057
2058 Value inY = arith::IndexCastOp::create(b, b.getI32Type(), y);
2059 Value inX = arith::IndexCastOp::create(b, b.getI32Type(), x);
2060
2061 SmallVector<int64_t> scale, offset, border;
2062 if (!tosa::getConstShapeValues(op.getScale().getDefiningOp(), scale) ||
2063 !tosa::getConstShapeValues(op.getOffset().getDefiningOp(), offset) ||
2064 !tosa::getConstShapeValues(op.getBorder().getDefiningOp(), border)) {
2065 return rewriter.notifyMatchFailure(
2066 op, "tosa.resize scale/offset/border should have compile time "
2067 "constant values.");
2068 }
2069
2070 Value yScaleN, yScaleD, xScaleN, xScaleD;
2071 yScaleN = arith::ConstantOp::create(b, b.getI32IntegerAttr(scale[0]));
2072 yScaleD = arith::ConstantOp::create(b, b.getI32IntegerAttr(scale[1]));
2073 xScaleN = arith::ConstantOp::create(b, b.getI32IntegerAttr(scale[2]));
2074 xScaleD = arith::ConstantOp::create(b, b.getI32IntegerAttr(scale[3]));
2075
2076 Value yOffset, xOffset, yBorder, xBorder;
2077 yOffset = arith::ConstantOp::create(b, b.getI32IntegerAttr(offset[0]));
2078 xOffset = arith::ConstantOp::create(b, b.getI32IntegerAttr(offset[1]));
2079 yBorder = arith::ConstantOp::create(b, b.getI32IntegerAttr(border[0]));
2080 xBorder = arith::ConstantOp::create(b, b.getI32IntegerAttr(border[1]));
2081
2082 // Compute the ix and dx values for both the X and Y dimensions.
2083 auto getIndexAndDeltaFp = [&](Value &index, Value &delta, Value in,
2084 Value scaleN, Value scaleD, Value offset,
2085 int size, ImplicitLocOpBuilder &b) {
2086 if (size == 1) {
2087 index = zeroI32;
2088 delta = zeroFp;
2089 return;
2090 }
2091 // x = x * scale_d + offset;
2092 // ix = floor(x / scale_n)
2093 Value val = arith::MulIOp::create(b, in, scaleD);
2094 val = arith::AddIOp::create(b, val, offset);
2095 index = arith::FloorDivSIOp::create(b, val, scaleN);
2096
2097 // rx = x - ix * scale_n (x % scale_n, if values are positive)
2098 Value scaledIndex = arith::MulIOp::create(b, index, scaleN);
2099 Value r = arith::SubIOp::create(b, val, scaledIndex);
2100 Value rFp = arith::SIToFPOp::create(b, floatTy, r);
2101
2102 // dx = rx / scale_n
2103 Value scaleNfp = arith::UIToFPOp::create(b, floatTy, scaleN);
2104 delta = arith::DivFOp::create(b, rFp, scaleNfp);
2105 };
2106
2107 // Compute the ix and dx values for the X and Y dimensions - int case.
2108 auto getIndexAndDeltaInt = [&](Value &index, Value &delta, Value in,
2109 Value scaleN, Value scaleD, Value offset,
2110 int size, ImplicitLocOpBuilder &b) {
2111 if (size == 1) {
2112 index = zeroI32;
2113 delta = zeroI32;
2114 return;
2115 }
2116 // x = x * scale_d + offset;
2117 // ix = floor(x / scale_n)
2118 // dx = x - ix * scale_n;
2119 Value val = arith::MulIOp::create(b, in, scaleD);
2120 val = arith::AddIOp::create(b, val, offset);
2121 index = arith::FloorDivSIOp::create(b, val, scaleN);
2122 delta = arith::MulIOp::create(b, index, scaleN);
2123 delta = arith::SubIOp::create(b, val, delta);
2124 };
2125
2126 Value ix, iy, dx, dy;
2127 if (floatingPointMode) {
2128 getIndexAndDeltaFp(iy, dy, inY, yScaleN, yScaleD, yOffset, imageH, b);
2129 getIndexAndDeltaFp(ix, dx, inX, xScaleN, xScaleD, xOffset, imageW, b);
2130 } else {
2131 getIndexAndDeltaInt(iy, dy, inY, yScaleN, yScaleD, yOffset, imageH, b);
2132 getIndexAndDeltaInt(ix, dx, inX, xScaleN, xScaleD, xOffset, imageW, b);
2133 }
2134
2135 if (op.getMode() == ResizeMode::NEAREST_NEIGHBOR) {
2136 auto one = arith::ConstantOp::create(b, b.getI32IntegerAttr(1));
2137
2138 auto getNearestIndexAndClamp = [&](Value val, Value dval, Value scale,
2139 Value max, int size,
2140 ImplicitLocOpBuilder &b) -> Value {
2141 if (size == 1) {
2143 }
2144
2145 Value pred;
2146 if (floatingPointMode) {
2147 auto h =
2148 arith::ConstantOp::create(b, b.getFloatAttr(floatTy, 0.5f));
2149 pred = arith::CmpFOp::create(b, arith::CmpFPredicate::OGE, dval, h);
2150 } else {
2151 Value dvalDouble = arith::ShLIOp::create(b, dval, one);
2152 pred = arith::CmpIOp::create(b, arith::CmpIPredicate::sge,
2153 dvalDouble, scale);
2154 }
2155
2156 auto offset = arith::SelectOp::create(b, pred, one, zeroI32);
2157 val = arith::AddIOp::create(b, val, offset);
2158 val = clampIntHelper(loc, val, zeroI32, max, b, /*isUnsigned=*/false);
2159 return arith::IndexCastOp::create(b, b.getIndexType(), val);
2160 };
2161
2162 iy = getNearestIndexAndClamp(iy, dy, yScaleN, hMax, imageH, b);
2163 ix = getNearestIndexAndClamp(ix, dx, xScaleN, wMax, imageW, b);
2164
2165 Value result = tensor::ExtractOp::create(
2166 b, input, ValueRange{batch, iy, ix, channel});
2167
2168 linalg::YieldOp::create(b, result);
2169 } else {
2170 // The mode here must be BILINEAR.
2171 assert(op.getMode() == ResizeMode::BILINEAR);
2172
2173 auto oneVal = arith::ConstantOp::create(b, b.getI32IntegerAttr(1));
2174
2175 auto getClampedIdxs = [&](Value &val0, Value &val1, int size, Value in,
2176 Value max, ImplicitLocOpBuilder &b) {
2177 val0 = in;
2178 val1 = arith::AddIOp::create(b, val0, oneVal);
2179 val0 =
2180 clampIntHelper(loc, val0, zeroI32, max, b, /*isUnsigned=*/false);
2181 val1 =
2182 clampIntHelper(loc, val1, zeroI32, max, b, /*isUnsigned=*/false);
2183 val0 = arith::IndexCastOp::create(b, b.getIndexType(), val0);
2184 val1 = arith::IndexCastOp::create(b, b.getIndexType(), val1);
2185 };
2186
2187 // Linalg equivalent to the section below:
2188 // int16_t iy0 = apply_max(iy, 0);
2189 // int16_t iy1 = apply_min(iy + 1, IH - 1);
2190 // int16_t ix0 = apply_max(ix, 0);
2191 // int16_t ix1 = apply_min(ix + 1, IW - 1);
2192 Value x0, x1, y0, y1;
2193 getClampedIdxs(y0, y1, imageH, iy, hMax, b);
2194 getClampedIdxs(x0, x1, imageW, ix, wMax, b);
2195
2196 Value y0x0 = tensor::ExtractOp::create(
2197 b, input, ValueRange{batch, y0, x0, channel});
2198 Value y0x1 = tensor::ExtractOp::create(
2199 b, input, ValueRange{batch, y0, x1, channel});
2200 Value y1x0 = tensor::ExtractOp::create(
2201 b, input, ValueRange{batch, y1, x0, channel});
2202 Value y1x1 = tensor::ExtractOp::create(
2203 b, input, ValueRange{batch, y1, x1, channel});
2204
2205 if (floatingPointMode) {
2206 auto oneVal =
2207 arith::ConstantOp::create(b, b.getFloatAttr(floatTy, 1.0f));
2208 auto interpolate = [&](Value val0, Value val1, Value delta,
2209 int inputSize,
2210 ImplicitLocOpBuilder &b) -> Value {
2211 if (inputSize == 1)
2212 return val0;
2213 Value oneMinusDelta = arith::SubFOp::create(b, oneVal, delta);
2214 Value mul0 = arith::MulFOp::create(b, val0, oneMinusDelta);
2215 Value mul1 = arith::MulFOp::create(b, val1, delta);
2216 return arith::AddFOp::create(b, mul0, mul1);
2217 };
2218
2219 // Linalg equivalent to the section below:
2220 // topAcc = v00 * (unit_x - dx);
2221 // topAcc += v01 * dx;
2222 Value topAcc = interpolate(y0x0, y0x1, dx, imageW, b);
2223
2224 // Linalg equivalent to the section below:
2225 // bottomAcc = v10 * (unit_x - dx);
2226 // bottomAcc += v11 * dx;
2227 Value bottomAcc = interpolate(y1x0, y1x1, dx, imageW, b);
2228
2229 // Linalg equivalent to the section below:
2230 // result = topAcc * (unit_y - dy) + bottomAcc * dy
2231 Value result = interpolate(topAcc, bottomAcc, dy, imageH, b);
2232 linalg::YieldOp::create(b, result);
2233 } else {
2234 // Perform in quantized space.
2235 y0x0 = arith::ExtSIOp::create(b, resultETy, y0x0);
2236 y0x1 = arith::ExtSIOp::create(b, resultETy, y0x1);
2237 y1x0 = arith::ExtSIOp::create(b, resultETy, y1x0);
2238 y1x1 = arith::ExtSIOp::create(b, resultETy, y1x1);
2239
2240 const int64_t deltaBitwidth = dx.getType().getIntOrFloatBitWidth();
2241 if (resultETy.getIntOrFloatBitWidth() > deltaBitwidth) {
2242 dx = arith::ExtSIOp::create(b, resultETy, dx);
2243 dy = arith::ExtSIOp::create(b, resultETy, dy);
2244 }
2245
2246 Value yScaleNExt = yScaleN;
2247 Value xScaleNExt = xScaleN;
2248
2249 const int64_t scaleBitwidth =
2250 xScaleN.getType().getIntOrFloatBitWidth();
2251 if (resultETy.getIntOrFloatBitWidth() > scaleBitwidth) {
2252 yScaleNExt = arith::ExtSIOp::create(b, resultETy, yScaleN);
2253 xScaleNExt = arith::ExtSIOp::create(b, resultETy, xScaleN);
2254 }
2255
2256 auto interpolate = [](Value val0, Value val1, Value weight1,
2257 Value scale, int inputSize,
2258 ImplicitLocOpBuilder &b) -> Value {
2259 if (inputSize == 1)
2260 return arith::MulIOp::create(b, val0, scale);
2261 Value weight0 = arith::SubIOp::create(b, scale, weight1);
2262 Value mul0 = arith::MulIOp::create(b, val0, weight0);
2263 Value mul1 = arith::MulIOp::create(b, val1, weight1);
2264 return arith::AddIOp::create(b, mul0, mul1);
2265 };
2266
2267 Value topAcc = interpolate(y0x0, y0x1, dx, xScaleNExt, imageW, b);
2268 Value bottomAcc = interpolate(y1x0, y1x1, dx, xScaleNExt, imageW, b);
2269 Value result =
2270 interpolate(topAcc, bottomAcc, dy, yScaleNExt, imageH, b);
2271 linalg::YieldOp::create(b, result);
2272 }
2273 }
2274 }
2275
2276 rewriter.replaceOp(op, resize);
2277 return success();
2278 }
2279};
2280
2281// At the codegen level any identity operations should be removed. Any cases
2282// where identity is load-bearing (e.g. cross device computation) should be
2283// handled before lowering to codegen.
2284template <typename SrcOp>
2285class IdentityNConverter : public OpRewritePattern<SrcOp> {
2286public:
2287 using OpRewritePattern<SrcOp>::OpRewritePattern;
2288
2289 LogicalResult matchAndRewrite(SrcOp op,
2290 PatternRewriter &rewriter) const final {
2291 rewriter.replaceOp(op, op.getOperation()->getOperands());
2292 return success();
2293 }
2294};
2295
2296template <typename SrcOp>
2297class ReduceConverter : public OpRewritePattern<SrcOp> {
2298public:
2299 ReduceConverter(MLIRContext *context, bool allowNonFinites)
2300 : OpRewritePattern<SrcOp>(context), allowNonFinites(allowNonFinites) {}
2301
2302 LogicalResult matchAndRewrite(SrcOp reduceOp,
2303 PatternRewriter &rewriter) const final {
2304 return reduceMatchAndRewriteHelper(reduceOp, reduceOp.getAxis(), rewriter,
2305 allowNonFinites);
2306 }
2307
2308private:
2309 bool allowNonFinites;
2310};
2311
2312class ReverseConverter : public OpRewritePattern<tosa::ReverseOp> {
2313public:
2314 using OpRewritePattern<tosa::ReverseOp>::OpRewritePattern;
2315
2316 LogicalResult matchAndRewrite(tosa::ReverseOp op,
2317 PatternRewriter &rewriter) const final {
2318 auto loc = op.getLoc();
2319 Value input = op.getInput1();
2320 auto inputTy = cast<ShapedType>(input.getType());
2321 auto resultTy = cast<ShapedType>(op.getType());
2322 auto axis = op.getAxis();
2323
2324 SmallVector<Value> dynDims;
2325 for (int i = 0; i < inputTy.getRank(); i++) {
2326 if (inputTy.isDynamicDim(i)) {
2327 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2328 }
2329 }
2330
2331 Value axisDimSize = tensor::DimOp::create(rewriter, loc, input, axis);
2332
2333 // First fill the output buffer with the init value.
2334 auto emptyTensor = tensor::EmptyOp::create(
2335 rewriter, loc, inputTy.getShape(),
2336 inputTy.getElementType(), ArrayRef<Value>({dynDims}))
2337 .getResult();
2338 SmallVector<AffineMap, 2> affineMaps = {
2339 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2340
2341 rewriter.replaceOpWithNewOp<linalg::GenericOp>(
2342 op, resultTy, ArrayRef<Value>({}), ValueRange{emptyTensor}, affineMaps,
2343 getNParallelLoopsAttrs(resultTy.getRank()),
2344 [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange args) {
2345 llvm::SmallVector<Value> indices;
2346 for (unsigned int i = 0; i < inputTy.getRank(); i++) {
2347 Value index =
2348 linalg::IndexOp::create(rewriter, nestedLoc, i).getResult();
2349 if (i == axis) {
2350 auto one = arith::ConstantIndexOp::create(rewriter, nestedLoc, 1);
2351 auto sizeMinusOne =
2352 arith::SubIOp::create(rewriter, nestedLoc, axisDimSize, one);
2353 index = arith::SubIOp::create(rewriter, nestedLoc, sizeMinusOne,
2354 index);
2355 }
2356
2357 indices.push_back(index);
2358 }
2359
2360 auto extract = tensor::ExtractOp::create(nestedBuilder, nestedLoc,
2361 input, indices);
2362 linalg::YieldOp::create(nestedBuilder, op.getLoc(),
2363 extract.getResult());
2364 });
2365 return success();
2366 }
2367};
2368
2369// This converter translate a tile operation to a reshape, broadcast, reshape.
2370// The first reshape minimally expands each tiled dimension to include a
2371// proceding size-1 dim. This dim is then broadcasted to the appropriate
2372// multiple.
2373struct TileConverter : public OpConversionPattern<tosa::TileOp> {
2374 using OpConversionPattern<tosa::TileOp>::OpConversionPattern;
2375
2376 LogicalResult
2377 matchAndRewrite(tosa::TileOp op, OpAdaptor adaptor,
2378 ConversionPatternRewriter &rewriter) const override {
2379 auto loc = op.getLoc();
2380 auto input = op.getInput1();
2381 auto inputTy = cast<ShapedType>(input.getType());
2382 auto inputShape = inputTy.getShape();
2383 auto resultTy = cast<ShapedType>(op.getType());
2384 auto elementTy = inputTy.getElementType();
2385 int64_t rank = inputTy.getRank();
2386
2387 SmallVector<int64_t> multiples;
2388 if (failed(op.getConstantMultiples(multiples)))
2389 return failure();
2390
2391 // Broadcast the newly added dimensions to their appropriate multiple.
2392 SmallVector<int64_t, 2> genericShape;
2393 for (int i = 0; i < rank; i++) {
2394 int64_t dim = multiples[i];
2395 genericShape.push_back(dim == -1 ? ShapedType::kDynamic : dim);
2396 genericShape.push_back(inputShape[i]);
2397 }
2398
2399 SmallVector<Value> dynDims;
2400 for (int i = 0; i < inputTy.getRank(); i++) {
2401 if (inputTy.isDynamicDim(i) || multiples[i] == -1) {
2402 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2403 }
2404 }
2405
2406 auto emptyTensor = tensor::EmptyOp::create(
2407 rewriter, op.getLoc(), genericShape, elementTy, dynDims);
2408
2409 // We needs to map the input shape to the non-broadcasted dimensions.
2410 SmallVector<AffineExpr, 4> dimExprs;
2411 dimExprs.reserve(rank);
2412 for (unsigned i = 0; i < rank; ++i)
2413 dimExprs.push_back(rewriter.getAffineDimExpr(i * 2 + 1));
2414
2415 auto readAffineMap =
2416 AffineMap::get(/*dimCount=*/rank * 2, /*symbolCount=*/0, dimExprs,
2417 rewriter.getContext());
2418
2419 SmallVector<AffineMap, 2> affineMaps = {
2420 readAffineMap, rewriter.getMultiDimIdentityMap(genericShape.size())};
2421
2422 auto genericOp = linalg::GenericOp::create(
2423 rewriter, loc, RankedTensorType::get(genericShape, elementTy), input,
2424 ValueRange{emptyTensor}, affineMaps,
2425 getNParallelLoopsAttrs(genericShape.size()),
2426 [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange args) {
2427 linalg::YieldOp::create(nestedBuilder, op.getLoc(), *args.begin());
2428 });
2429
2430 auto shapeValue = getTosaConstShape(
2431 rewriter, loc, mlir::tosa::convertFromMlirShape(resultTy.getShape()));
2432 rewriter.replaceOpWithNewOp<tosa::ReshapeOp>(
2433 op, resultTy, genericOp.getResult(0), shapeValue);
2434 return success();
2435 }
2436};
2437
2438// Tosa argmax lowering represents the ArgMax op as an linalg.indexed_generic
2439// op, producing two output buffers.
2440//
2441// The first output buffer contains the index of the found maximum value. It is
2442// initialized to 0 and is resulting integer type.
2443//
2444// The second output buffer contains the maximum value found. It is initialized
2445// to the minimum representable value of the input element type. After being
2446// populated by indexed_generic, this buffer is disgarded as only the index is
2447// requested.
2448//
2449// The indexed_generic op updates both the maximum value and index if the
2450// current value exceeds the running max.
2451class ArgMaxConverter : public OpRewritePattern<tosa::ArgMaxOp> {
2452public:
2453 ArgMaxConverter(MLIRContext *context, bool allowNonFinites)
2454 : OpRewritePattern<tosa::ArgMaxOp>(context),
2455 allowNonFinites(allowNonFinites) {}
2456
2457 LogicalResult matchAndRewrite(tosa::ArgMaxOp argmaxOp,
2458 PatternRewriter &rewriter) const final {
2459 auto loc = argmaxOp.getLoc();
2460 Value input = argmaxOp.getInput();
2461 auto inputTy = cast<ShapedType>(input.getType());
2462 auto resultTy = cast<ShapedType>(argmaxOp.getOutput().getType());
2463 auto inElementTy = inputTy.getElementType();
2464 auto outElementTy = resultTy.getElementType();
2465 int axis = argmaxOp.getAxis();
2466 auto resultMaxTy = RankedTensorType::get(resultTy.getShape(), inElementTy);
2467
2468 if (!isa<IntegerType>(outElementTy))
2469 return rewriter.notifyMatchFailure(
2470 argmaxOp,
2471 "tosa.arg_max to linalg.* requires integer-like result type");
2472
2473 SmallVector<Value> dynDims;
2474 for (int i = 0; i < inputTy.getRank(); i++) {
2475 if (inputTy.isDynamicDim(i) && i != axis) {
2476 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2477 }
2478 }
2479
2480 // First fill the output buffer for the index.
2481 auto emptyTensorIdx =
2482 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2483 outElementTy, dynDims)
2484 .getResult();
2485 auto fillValueIdx = arith::ConstantOp::create(
2486 rewriter, loc, rewriter.getIntegerAttr(outElementTy, 0));
2487 auto filledTensorIdx =
2488 linalg::FillOp::create(rewriter, loc, ValueRange{fillValueIdx},
2489 ValueRange{emptyTensorIdx})
2490 .result();
2491
2492 // Second fill the output buffer for the running max.
2493 auto emptyTensorMax =
2494 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(), inElementTy,
2495 dynDims)
2496 .getResult();
2497 auto fillValueMaxAttr = createInitialValueForReduceOp(
2498 argmaxOp, inElementTy, rewriter, allowNonFinites);
2499
2500 if (!fillValueMaxAttr)
2501 return rewriter.notifyMatchFailure(
2502 argmaxOp, "unsupported tosa.argmax element type");
2503
2504 auto fillValueMax =
2505 arith::ConstantOp::create(rewriter, loc, fillValueMaxAttr);
2506 auto filledTensorMax =
2507 linalg::FillOp::create(rewriter, loc, ValueRange{fillValueMax},
2508 ValueRange{emptyTensorMax})
2509 .result();
2510
2511 // We need to reduce along the arg-max axis, with parallel operations along
2512 // the rest.
2513 SmallVector<utils::IteratorType, 4> iteratorTypes;
2514 iteratorTypes.resize(inputTy.getRank(), utils::IteratorType::parallel);
2515 iteratorTypes[axis] = utils::IteratorType::reduction;
2516
2517 SmallVector<AffineExpr, 2> srcExprs;
2518 SmallVector<AffineExpr, 2> dstExprs;
2519 for (int i = 0, rank = inputTy.getRank(); i != rank; ++i) {
2520 srcExprs.push_back(mlir::getAffineDimExpr(i, rewriter.getContext()));
2521 if (axis != i)
2522 dstExprs.push_back(mlir::getAffineDimExpr(i, rewriter.getContext()));
2523 }
2524
2525 bool didEncounterError = false;
2526 auto maps = AffineMap::inferFromExprList({srcExprs, dstExprs, dstExprs},
2527 rewriter.getContext());
2528 auto linalgOp = linalg::GenericOp::create(
2529 rewriter, loc, ArrayRef<Type>({resultTy, resultMaxTy}), input,
2530 ValueRange({filledTensorIdx, filledTensorMax}), maps, iteratorTypes,
2531 [&](OpBuilder &nestedBuilder, Location nestedLoc,
2532 ValueRange blockArgs) {
2533 auto newValue = blockArgs[0];
2534 auto oldIndex = blockArgs[1];
2535 auto oldValue = blockArgs[2];
2536
2537 Value newIndex = arith::IndexCastOp::create(
2538 rewriter, nestedLoc, oldIndex.getType(),
2539 linalg::IndexOp::create(rewriter, loc, axis));
2540
2541 Value predicate;
2542 if (isa<FloatType>(inElementTy)) {
2543 if (argmaxOp.getNanMode() == NanPropagationMode::IGNORE) {
2544 // Only update index & max value for non NaN values. If all
2545 // values are NaNs, the initial index will be return which is 0.
2546 predicate = arith::CmpFOp::create(rewriter, nestedLoc,
2547 arith::CmpFPredicate::OGT,
2548 newValue, oldValue);
2549 } else {
2550 // Update max value if either of the following is true:
2551 // - new value is bigger
2552 // - cur max is not NaN and new value is NaN
2553 Value gt = arith::CmpFOp::create(rewriter, nestedLoc,
2554 arith::CmpFPredicate::UGT,
2555 newValue, oldValue);
2556 Value oldNonNaN = arith::CmpFOp::create(rewriter, nestedLoc,
2557 arith::CmpFPredicate::ORD,
2558 oldValue, oldValue);
2559 predicate = arith::AndIOp::create(
2560 rewriter, nestedLoc, rewriter.getI1Type(), gt, oldNonNaN);
2561 }
2562 } else if (isa<IntegerType>(inElementTy)) {
2563 predicate = arith::CmpIOp::create(rewriter, nestedLoc,
2564 arith::CmpIPredicate::sgt,
2565 newValue, oldValue);
2566 } else {
2567 didEncounterError = true;
2568 return;
2569 }
2570
2571 auto resultMax = arith::SelectOp::create(
2572 rewriter, nestedLoc, predicate, newValue, oldValue);
2573 auto resultIndex = arith::SelectOp::create(
2574 rewriter, nestedLoc, predicate, newIndex, oldIndex);
2575 linalg::YieldOp::create(nestedBuilder, nestedLoc,
2576 ValueRange({resultIndex, resultMax}));
2577 });
2578
2579 if (didEncounterError)
2580 return rewriter.notifyMatchFailure(
2581 argmaxOp, "unsupported tosa.argmax element type");
2582
2583 rewriter.replaceOp(argmaxOp, linalgOp.getResult(0));
2584 return success();
2585 }
2586
2587private:
2588 bool allowNonFinites;
2589};
2590
2591class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
2592public:
2593 using OpConversionPattern<tosa::GatherOp>::OpConversionPattern;
2594 LogicalResult
2595 matchAndRewrite(tosa::GatherOp op, OpAdaptor adaptor,
2596 ConversionPatternRewriter &rewriter) const final {
2597 auto input = adaptor.getOperands()[0];
2598 auto indices = adaptor.getOperands()[1];
2599
2600 auto valuesTy = dyn_cast<RankedTensorType>(op.getValues().getType());
2601 auto resultTy = dyn_cast<RankedTensorType>(op.getType());
2602 if (!valuesTy || !resultTy)
2603 return rewriter.notifyMatchFailure(op, "unranked tensors not supported");
2604
2605 auto dynamicDims = inferDynamicDimsForGather(
2606 rewriter, op.getLoc(), adaptor.getValues(), adaptor.getIndices());
2607
2608 auto resultElementTy = resultTy.getElementType();
2609
2610 auto loc = op.getLoc();
2611 auto emptyTensor =
2612 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2613 resultElementTy, dynamicDims)
2614 .getResult();
2615
2616 SmallVector<AffineMap, 2> affineMaps = {
2618 /*dimCount=*/resultTy.getRank(), /*symbolCount=*/0,
2619 {rewriter.getAffineDimExpr(0), rewriter.getAffineDimExpr(1)},
2620 rewriter.getContext()),
2621 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2622
2623 auto genericOp = linalg::GenericOp::create(
2624 rewriter, loc, ArrayRef<Type>({resultTy}), ValueRange{indices},
2625 ValueRange{emptyTensor}, affineMaps,
2626 getNParallelLoopsAttrs(resultTy.getRank()),
2627 [&](OpBuilder &b, Location loc, ValueRange args) {
2628 auto indexValue = args[0];
2629 auto index0 = linalg::IndexOp::create(rewriter, loc, 0);
2630 Value index1 = arith::IndexCastOp::create(
2631 rewriter, loc, rewriter.getIndexType(), indexValue);
2632 auto index2 = linalg::IndexOp::create(rewriter, loc, 2);
2633 Value extract = tensor::ExtractOp::create(
2634 rewriter, loc, input, ValueRange{index0, index1, index2});
2635 linalg::YieldOp::create(rewriter, loc, extract);
2636 });
2637 rewriter.replaceOp(op, genericOp.getResult(0));
2638 return success();
2639 }
2640
2641 static llvm::SmallVector<Value> inferDynamicDimsForGather(OpBuilder &builder,
2642 Location loc,
2643 Value values,
2644 Value indices) {
2645 llvm::SmallVector<Value> results;
2646
2647 auto addDynamicDimension = [&](Value source, int64_t dim) {
2648 auto sz = tensor::getMixedSize(builder, loc, source, dim);
2649 if (auto dimValue = llvm::dyn_cast_if_present<Value>(sz))
2650 results.push_back(dimValue);
2651 };
2652
2653 addDynamicDimension(values, 0);
2654 addDynamicDimension(indices, 1);
2655 addDynamicDimension(values, 2);
2656 return results;
2657 }
2658};
2659
2660// Lowerings the TableOp to a series of gathers and numerica operations. This
2661// includes interpolation between the high/low values. For the I8 varient, this
2662// simplifies to a single gather operation.
2663class TableConverter : public OpRewritePattern<tosa::TableOp> {
2664public:
2665 using OpRewritePattern<tosa::TableOp>::OpRewritePattern;
2666
2667 LogicalResult matchAndRewrite(tosa::TableOp op,
2668 PatternRewriter &rewriter) const final {
2669 auto loc = op.getLoc();
2670 Value input = op.getInput1();
2671 Value table = op.getTable();
2672 auto inputTy = cast<ShapedType>(input.getType());
2673 auto tableTy = cast<ShapedType>(table.getType());
2674 auto resultTy = cast<ShapedType>(op.getType());
2675
2676 auto inputElementTy = inputTy.getElementType();
2677 auto tableElementTy = tableTy.getElementType();
2678 auto resultElementTy = resultTy.getElementType();
2679
2680 SmallVector<Value> dynDims;
2681 for (int i = 0; i < resultTy.getRank(); ++i) {
2682 if (inputTy.isDynamicDim(i)) {
2683 dynDims.push_back(
2684 tensor::DimOp::create(rewriter, loc, op.getOperand(0), i));
2685 }
2686 }
2687
2688 auto emptyTensor =
2689 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2690 resultElementTy, dynDims)
2691 .getResult();
2692
2693 SmallVector<AffineMap, 2> affineMaps = {
2694 rewriter.getMultiDimIdentityMap(resultTy.getRank()),
2695 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2696
2697 auto genericOp = linalg::GenericOp::create(
2698 rewriter, loc, resultTy, ValueRange({input}), ValueRange{emptyTensor},
2699 affineMaps, getNParallelLoopsAttrs(resultTy.getRank()));
2700 rewriter.replaceOp(op, genericOp.getResult(0));
2701
2702 {
2703 OpBuilder::InsertionGuard regionGuard(rewriter);
2704 Block *block = rewriter.createBlock(
2705 &genericOp.getRegion(), genericOp.getRegion().end(),
2706 TypeRange({inputElementTy, resultElementTy}), {loc, loc});
2707
2708 auto inputValue = block->getArgument(0);
2709 rewriter.setInsertionPointToStart(block);
2710 if (inputElementTy.isInteger(8) && tableElementTy.isInteger(8) &&
2711 resultElementTy.isInteger(8)) {
2712 Value index = arith::IndexCastOp::create(
2713 rewriter, loc, rewriter.getIndexType(), inputValue);
2714 Value offset = arith::ConstantIndexOp::create(rewriter, loc, 128);
2715 index = arith::AddIOp::create(rewriter, loc, rewriter.getIndexType(),
2716 index, offset);
2717 Value extract =
2718 tensor::ExtractOp::create(rewriter, loc, table, ValueRange{index});
2719 linalg::YieldOp::create(rewriter, loc, extract);
2720 return success();
2721 }
2722
2723 if (inputElementTy.isInteger(16) && tableElementTy.isInteger(16) &&
2724 resultElementTy.isInteger(32)) {
2725 Value extend = arith::ExtSIOp::create(
2726 rewriter, loc, rewriter.getI32Type(), inputValue);
2727
2728 auto offset = arith::ConstantOp::create(
2729 rewriter, loc, rewriter.getI32IntegerAttr(32768));
2730 auto seven = arith::ConstantOp::create(rewriter, loc,
2731 rewriter.getI32IntegerAttr(7));
2732 auto one = arith::ConstantOp::create(rewriter, loc,
2733 rewriter.getI32IntegerAttr(1));
2734 auto b1111111 = arith::ConstantOp::create(
2735 rewriter, loc, rewriter.getI32IntegerAttr(127));
2736
2737 // Compute the index and fractional part from the input value:
2738 // value = value + 32768
2739 // index = value >> 7;
2740 // fraction = 0x01111111 & value
2741 auto extendAdd = arith::AddIOp::create(rewriter, loc, extend, offset);
2742 Value index = arith::ShRUIOp::create(rewriter, loc, extendAdd, seven);
2743 Value fraction =
2744 arith::AndIOp::create(rewriter, loc, extendAdd, b1111111);
2745
2746 // Extract the base and next values from the table.
2747 // base = (int32_t) table[index];
2748 // next = (int32_t) table[index + 1];
2749 Value indexPlusOne = arith::AddIOp::create(rewriter, loc, index, one);
2750
2751 index = arith::IndexCastOp::create(rewriter, loc,
2752 rewriter.getIndexType(), index);
2753 indexPlusOne = arith::IndexCastOp::create(
2754 rewriter, loc, rewriter.getIndexType(), indexPlusOne);
2755
2756 Value base =
2757 tensor::ExtractOp::create(rewriter, loc, table, ValueRange{index});
2758 Value next = tensor::ExtractOp::create(rewriter, loc, table,
2759 ValueRange{indexPlusOne});
2760
2761 base =
2762 arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), base);
2763 next =
2764 arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), next);
2765
2766 // Use the fractional part to interpolate between the input values:
2767 // result = (base << 7) + (next - base) * fraction
2768 Value baseScaled = arith::ShLIOp::create(rewriter, loc, base, seven);
2769 Value diff = arith::SubIOp::create(rewriter, loc, next, base);
2770 Value diffScaled = arith::MulIOp::create(rewriter, loc, diff, fraction);
2771 Value result =
2772 arith::AddIOp::create(rewriter, loc, baseScaled, diffScaled);
2773
2774 linalg::YieldOp::create(rewriter, loc, result);
2775
2776 return success();
2777 }
2778 }
2779
2780 return rewriter.notifyMatchFailure(
2781 op, "unable to create body for tosa.table op");
2782 }
2783};
2784
2785struct RFFT2dConverter final : public OpRewritePattern<RFFT2dOp> {
2786 using OpRewritePattern<RFFT2dOp>::OpRewritePattern;
2787
2788 static bool isRankedTensor(Type type) { return isa<RankedTensorType>(type); }
2789
2790 static OpFoldResult halfPlusOne(OpBuilder &builder, Location loc,
2791 OpFoldResult ofr) {
2792 auto one = arith::ConstantIndexOp::create(builder, loc, 1);
2793 auto two = arith::ConstantIndexOp::create(builder, loc, 2);
2794
2795 auto value = getValueOrCreateConstantIndexOp(builder, loc, ofr);
2796 auto divBy2 = builder.createOrFold<arith::DivUIOp>(loc, value, two);
2797 auto plusOne = builder.createOrFold<arith::AddIOp>(loc, divBy2, one);
2798 return getAsOpFoldResult(plusOne);
2799 }
2800
2801 static RankedTensorType
2802 computeOutputShape(OpBuilder &builder, Location loc, Value input,
2803 llvm::SmallVectorImpl<Value> &dynamicSizes) {
2804 // Get [N, H, W]
2805 auto dims = tensor::getMixedSizes(builder, loc, input);
2806
2807 // Set W = (W / 2) + 1 to account for the half-sized W dimension of the
2808 // output tensors.
2809 dims[2] = halfPlusOne(builder, loc, dims[2]);
2810
2811 llvm::SmallVector<int64_t, 3> staticSizes;
2812 dispatchIndexOpFoldResults(dims, dynamicSizes, staticSizes);
2813
2814 auto elementType = cast<RankedTensorType>(input.getType()).getElementType();
2815 return RankedTensorType::get(staticSizes, elementType);
2816 }
2817
2818 static Value createZeroTensor(PatternRewriter &rewriter, Location loc,
2819 RankedTensorType type,
2820 llvm::ArrayRef<Value> dynamicSizes) {
2821 auto emptyTensor =
2822 tensor::EmptyOp::create(rewriter, loc, type, dynamicSizes);
2823 auto fillValueAttr = rewriter.getZeroAttr(type.getElementType());
2824 auto fillValue = arith::ConstantOp::create(rewriter, loc, fillValueAttr);
2825 auto filledTensor =
2826 linalg::FillOp::create(rewriter, loc, ValueRange{fillValue},
2827 ValueRange{emptyTensor})
2828 .result();
2829 return filledTensor;
2830 }
2831
2832 static Value castIndexToFloat(OpBuilder &builder, Location loc,
2833 FloatType type, Value value) {
2834 auto integerVal = arith::IndexCastUIOp::create(
2835 builder, loc,
2836 type.getIntOrFloatBitWidth() > 32 ? builder.getI64Type()
2837 : builder.getI32Type(),
2838 value);
2839
2840 return arith::UIToFPOp::create(builder, loc, type, integerVal);
2841 }
2842
2843 static Value createLinalgIndex(OpBuilder &builder, Location loc,
2844 FloatType type, int64_t index) {
2845 auto indexVal = linalg::IndexOp::create(builder, loc, index);
2846 return castIndexToFloat(builder, loc, type, indexVal);
2847 }
2848
2849 template <typename... Args>
2850 static llvm::SmallVector<AffineExpr, 4> affineDimsExpr(OpBuilder &builder,
2851 Args... args) {
2852 return {builder.getAffineDimExpr(args)...};
2853 }
2854
2855 LogicalResult matchAndRewrite(RFFT2dOp rfft2d,
2856 PatternRewriter &rewriter) const override {
2857 if (!llvm::all_of(rfft2d->getOperandTypes(), isRankedTensor) ||
2858 !llvm::all_of(rfft2d->getResultTypes(), isRankedTensor)) {
2859 return rewriter.notifyMatchFailure(rfft2d,
2860 "only supports ranked tensors");
2861 }
2862
2863 auto loc = rfft2d.getLoc();
2864 auto input = rfft2d.getInputReal();
2865 auto elementType =
2866 dyn_cast<FloatType>(cast<ShapedType>(input.getType()).getElementType());
2867 if (!elementType)
2868 return rewriter.notifyMatchFailure(rfft2d,
2869 "only supports float element types");
2870
2871 // Compute the output type and set of dynamic sizes
2872 llvm::SmallVector<Value> dynamicSizes;
2873 auto outputType = computeOutputShape(rewriter, loc, input, dynamicSizes);
2874
2875 // Iterator types for the linalg.generic implementation
2876 llvm::SmallVector<utils::IteratorType, 5> iteratorTypes = {
2877 utils::IteratorType::parallel, utils::IteratorType::parallel,
2878 utils::IteratorType::parallel, utils::IteratorType::reduction,
2879 utils::IteratorType::reduction};
2880
2881 // Inputs/outputs to the linalg.generic implementation
2882 llvm::SmallVector<Value> genericOpInputs = {input};
2883 llvm::SmallVector<Value> genericOpOutputs = {
2884 createZeroTensor(rewriter, loc, outputType, dynamicSizes),
2885 createZeroTensor(rewriter, loc, outputType, dynamicSizes)};
2886
2887 // Indexing maps for input and output tensors
2888 auto indexingMaps = AffineMap::inferFromExprList(
2889 llvm::ArrayRef{affineDimsExpr(rewriter, 0, 3, 4),
2890 affineDimsExpr(rewriter, 0, 1, 2),
2891 affineDimsExpr(rewriter, 0, 1, 2)},
2892 rewriter.getContext());
2893
2894 // Width and height dimensions of the original input.
2895 auto dimH = rewriter.createOrFold<tensor::DimOp>(loc, input, 1);
2896 auto dimW = rewriter.createOrFold<tensor::DimOp>(loc, input, 2);
2897
2898 // Constants and dimension sizes
2899 auto zeroFloat = arith::ConstantOp::create(
2900 rewriter, loc, rewriter.getZeroAttr(elementType));
2901 auto twoPiAttr = rewriter.getFloatAttr(elementType, 6.283185307179586);
2902 auto twoPi = arith::ConstantOp::create(rewriter, loc, twoPiAttr);
2903
2904 auto zeroIndex = arith::ConstantIndexOp::create(rewriter, loc, 0);
2905 auto twoIndex = arith::ConstantIndexOp::create(rewriter, loc, 2);
2906
2907 auto constH = castIndexToFloat(rewriter, loc, elementType, dimH);
2908 auto constW = castIndexToFloat(rewriter, loc, elementType, dimW);
2909 auto halfH = index::DivUOp::create(rewriter, loc, dimH, twoIndex);
2910 auto halfW = index::DivUOp::create(rewriter, loc, dimW, twoIndex);
2911
2912 auto buildBody = [&](OpBuilder &builder, Location loc, ValueRange args) {
2913 Value valReal = args[0];
2914 Value sumReal = args[1];
2915 Value sumImag = args[2];
2916
2917 // Indices for angle computation
2918 Value oy = linalg::IndexOp::create(builder, loc, 1);
2919 Value ox = linalg::IndexOp::create(builder, loc, 2);
2920 Value iy = linalg::IndexOp::create(builder, loc, 3);
2921 Value ix = linalg::IndexOp::create(builder, loc, 4);
2922
2923 // Calculating angle without integer parts of components as sin/cos are
2924 // periodic: angle = 2 * pi() * ( ( (iy * oy) % H) / H + ( (ix * ox) % W )
2925 // / W);
2926 auto iyXoy = index::MulOp::create(builder, loc, iy, oy);
2927 auto ixXox = index::MulOp::create(builder, loc, ix, ox);
2928
2929 auto iyRem = index::RemUOp::create(builder, loc, iyXoy, dimH);
2930 auto ixRem = index::RemUOp::create(builder, loc, ixXox, dimW);
2931
2932 auto iyRemFloat = castIndexToFloat(builder, loc, elementType, iyRem);
2933 auto ixRemFloat = castIndexToFloat(builder, loc, elementType, ixRem);
2934
2935 auto yComponent = arith::DivFOp::create(builder, loc, iyRemFloat, constH);
2936 auto xComponent = arith::DivFOp::create(builder, loc, ixRemFloat, constW);
2937 auto sumXY = arith::AddFOp::create(builder, loc, yComponent, xComponent);
2938 auto angle = arith::MulFOp::create(builder, loc, twoPi, sumXY);
2939
2940 // We will check the indices to see if this is a position that should use
2941 // a 0.0 weight for the imaginary value computation following the TOSA
2942 // specification with `tosa_extra_multiplies=true`.
2943 //
2944 // These are the relevant locations: (0,0), (0,W/2), (H/2,0), (H/2, W/2).
2945 auto iyIs0 = arith::CmpIOp::create(builder, loc, arith::CmpIPredicate::eq,
2946 iyRem, zeroIndex);
2947 auto iyIsHalfH = arith::CmpIOp::create(
2948 builder, loc, arith::CmpIPredicate::eq, iyRem, halfH);
2949 auto ixIs0 = arith::CmpIOp::create(builder, loc, arith::CmpIPredicate::eq,
2950 ixRem, zeroIndex);
2951 auto ixIsHalfW = arith::CmpIOp::create(
2952 builder, loc, arith::CmpIPredicate::eq, ixRem, halfW);
2953
2954 auto iyIsSinSkippable =
2955 arith::OrIOp::create(builder, loc, iyIs0, iyIsHalfH);
2956 auto ixIsSinSkippable =
2957 arith::OrIOp::create(builder, loc, ixIs0, ixIsHalfW);
2958 auto shouldSkipSin = arith::AndIOp::create(builder, loc, iyIsSinSkippable,
2959 ixIsSinSkippable);
2960
2961 // realComponent = valReal * cos(angle)
2962 // imagComponent = valReal * (shouldSkipSin ? 0.0 : sin(angle))
2963 auto cosAngle = math::CosOp::create(builder, loc, angle);
2964 auto sinAngle = math::SinOp::create(builder, loc, angle);
2965 auto imagWeight = arith::SelectOp::create(builder, loc, shouldSkipSin,
2966 zeroFloat, sinAngle);
2967 auto realComponent =
2968 arith::MulFOp::create(builder, loc, valReal, cosAngle);
2969 auto imagComponent =
2970 arith::MulFOp::create(builder, loc, valReal, imagWeight);
2971
2972 // outReal = sumReal + realComponent
2973 // outImag = sumImag - imagComponent
2974 auto outReal =
2975 arith::AddFOp::create(builder, loc, sumReal, realComponent);
2976 auto outImag =
2977 arith::SubFOp::create(builder, loc, sumImag, imagComponent);
2978
2979 linalg::YieldOp::create(builder, loc, ValueRange{outReal, outImag});
2980 };
2981
2982 rewriter.replaceOpWithNewOp<linalg::GenericOp>(
2983 rfft2d, rfft2d.getResultTypes(), genericOpInputs, genericOpOutputs,
2984 indexingMaps, iteratorTypes, buildBody);
2985
2986 return success();
2987 }
2988};
2989
2990struct FFT2dConverter final : OpRewritePattern<FFT2dOp> {
2992
2993 LogicalResult matchAndRewrite(FFT2dOp fft2d,
2994 PatternRewriter &rewriter) const override {
2995 if (!llvm::all_of(fft2d->getOperandTypes(),
2996 RFFT2dConverter::isRankedTensor) ||
2997 !llvm::all_of(fft2d->getResultTypes(),
2998 RFFT2dConverter::isRankedTensor)) {
2999 return rewriter.notifyMatchFailure(fft2d, "only supports ranked tensors");
3000 }
3001
3002 Location loc = fft2d.getLoc();
3003 Value input_real = fft2d.getInputReal();
3004 Value input_imag = fft2d.getInputImag();
3005 BoolAttr inverse = fft2d.getInverseAttr();
3006
3007 auto real_el_ty = cast<FloatType>(
3008 cast<ShapedType>(input_real.getType()).getElementType());
3009 [[maybe_unused]] auto imag_el_ty = cast<FloatType>(
3010 cast<ShapedType>(input_imag.getType()).getElementType());
3011
3012 assert(real_el_ty == imag_el_ty);
3013
3014 // Compute the output type and set of dynamic sizes
3015 SmallVector<Value> dynamicSizes;
3016
3017 // Get [N, H, W]
3018 auto dims = tensor::getMixedSizes(rewriter, loc, input_real);
3019
3020 SmallVector<int64_t, 3> staticSizes;
3021 dispatchIndexOpFoldResults(dims, dynamicSizes, staticSizes);
3022
3023 auto outputType = RankedTensorType::get(staticSizes, real_el_ty);
3024
3025 // Iterator types for the linalg.generic implementation
3026 SmallVector<utils::IteratorType, 5> iteratorTypes = {
3027 utils::IteratorType::parallel, utils::IteratorType::parallel,
3028 utils::IteratorType::parallel, utils::IteratorType::reduction,
3029 utils::IteratorType::reduction};
3030
3031 // Inputs/outputs to the linalg.generic implementation
3032 SmallVector<Value> genericOpInputs = {input_real, input_imag};
3033 SmallVector<Value> genericOpOutputs = {
3034 RFFT2dConverter::createZeroTensor(rewriter, loc, outputType,
3035 dynamicSizes),
3036 RFFT2dConverter::createZeroTensor(rewriter, loc, outputType,
3037 dynamicSizes)};
3038
3039 // Indexing maps for input and output tensors
3040 auto indexingMaps = AffineMap::inferFromExprList(
3041 ArrayRef{RFFT2dConverter::affineDimsExpr(rewriter, 0, 3, 4),
3042 RFFT2dConverter::affineDimsExpr(rewriter, 0, 3, 4),
3043 RFFT2dConverter::affineDimsExpr(rewriter, 0, 1, 2),
3044 RFFT2dConverter::affineDimsExpr(rewriter, 0, 1, 2)},
3045 rewriter.getContext());
3046
3047 // Width and height dimensions of the original input.
3048 auto dimH = rewriter.createOrFold<tensor::DimOp>(loc, input_real, 1);
3049 auto dimW = rewriter.createOrFold<tensor::DimOp>(loc, input_real, 2);
3050
3051 // Constants and dimension sizes
3052 auto twoPiAttr = rewriter.getFloatAttr(real_el_ty, 6.283185307179586);
3053 auto twoPi = arith::ConstantOp::create(rewriter, loc, twoPiAttr);
3054 Value constH =
3055 RFFT2dConverter::castIndexToFloat(rewriter, loc, real_el_ty, dimH);
3056 Value constW =
3057 RFFT2dConverter::castIndexToFloat(rewriter, loc, real_el_ty, dimW);
3058
3059 auto buildBody = [&](OpBuilder &builder, Location loc, ValueRange args) {
3060 Value valReal = args[0];
3061 Value valImag = args[1];
3062 Value sumReal = args[2];
3063 Value sumImag = args[3];
3064
3065 // Indices for angle computation
3066 Value oy = linalg::IndexOp::create(builder, loc, 1);
3067 Value ox = linalg::IndexOp::create(builder, loc, 2);
3068 Value iy = linalg::IndexOp::create(builder, loc, 3);
3069 Value ix = linalg::IndexOp::create(builder, loc, 4);
3070
3071 // float_t angle = sign_val * 2 * pi() * ( ( (iy * oy) % H) / H + ( (ix *
3072 // ox) % W ) / W);
3073 auto iyXoy = index::MulOp::create(builder, loc, iy, oy);
3074 auto ixXox = index::MulOp::create(builder, loc, ix, ox);
3075
3076 auto iyRem = index::RemUOp::create(builder, loc, iyXoy, dimH);
3077 auto ixRem = index::RemUOp::create(builder, loc, ixXox, dimW);
3078
3079 auto iyRemFloat =
3080 RFFT2dConverter::castIndexToFloat(builder, loc, real_el_ty, iyRem);
3081 auto ixRemFloat =
3082 RFFT2dConverter::castIndexToFloat(builder, loc, real_el_ty, ixRem);
3083
3084 auto yComponent = arith::DivFOp::create(builder, loc, iyRemFloat, constH);
3085 auto xComponent = arith::DivFOp::create(builder, loc, ixRemFloat, constW);
3086
3087 auto sumXY = arith::AddFOp::create(builder, loc, yComponent, xComponent);
3088 auto angle = arith::MulFOp::create(builder, loc, twoPi, sumXY);
3089
3090 if (inverse.getValue()) {
3091 angle = arith::MulFOp::create(
3092 builder, loc, angle,
3093 arith::ConstantOp::create(rewriter, loc,
3094 rewriter.getFloatAttr(real_el_ty, -1.0)));
3095 }
3096
3097 // realComponent = val_real * cos(a) + val_imag * sin(a);
3098 // imagComponent = -val_real * sin(a) + val_imag * cos(a);
3099 auto cosAngle = math::CosOp::create(builder, loc, angle);
3100 auto sinAngle = math::SinOp::create(builder, loc, angle);
3101
3102 auto rcos = arith::MulFOp::create(builder, loc, valReal, cosAngle);
3103 auto rsin = arith::MulFOp::create(builder, loc, valImag, sinAngle);
3104 auto realComponent = arith::AddFOp::create(builder, loc, rcos, rsin);
3105
3106 auto icos = arith::MulFOp::create(builder, loc, valImag, cosAngle);
3107 auto isin = arith::MulFOp::create(builder, loc, valReal, sinAngle);
3108
3109 auto imagComponent = arith::SubFOp::create(builder, loc, icos, isin);
3110
3111 // outReal = sumReal + realComponent
3112 // outImag = sumImag - imagComponent
3113 auto outReal =
3114 arith::AddFOp::create(builder, loc, sumReal, realComponent);
3115 auto outImag =
3116 arith::AddFOp::create(builder, loc, sumImag, imagComponent);
3117
3118 linalg::YieldOp::create(builder, loc, ValueRange{outReal, outImag});
3119 };
3120
3121 rewriter.replaceOpWithNewOp<linalg::GenericOp>(
3122 fft2d, fft2d.getResultTypes(), genericOpInputs, genericOpOutputs,
3123 indexingMaps, iteratorTypes, buildBody);
3124
3125 return success();
3126 }
3127};
3128
3129} // namespace
3130
3132 const TypeConverter &converter, RewritePatternSet *patterns,
3133 const TosaToLinalgOptions &options) {
3134
3135 // We have multiple resize coverters to handle degenerate cases.
3136 patterns->add<GenericResizeConverter>(patterns->getContext(),
3137 /*benefit=*/100);
3138 patterns->add<ResizeUnaryConverter>(patterns->getContext(),
3139 /*benefit=*/200);
3140 patterns->add<MaterializeResizeBroadcast>(patterns->getContext(),
3141 /*benefit=*/300);
3142
3143 patterns->add<
3144 // clang-format off
3145 PointwiseConverter<tosa::AddOp>,
3146 PointwiseConverter<tosa::SubOp>,
3147 PointwiseConverter<tosa::MulOp>,
3148 PointwiseConverter<tosa::IntDivOp>,
3149 PointwiseConverter<tosa::NegateOp>,
3150 PointwiseConverter<tosa::PowOp>,
3151 PointwiseConverter<tosa::ReciprocalOp>,
3152 PointwiseConverter<tosa::RsqrtOp>,
3153 PointwiseConverter<tosa::LogOp>,
3154 PointwiseConverter<tosa::ExpOp>,
3155 PointwiseConverter<tosa::AbsOp>,
3156 PointwiseConverter<tosa::SinOp>,
3157 PointwiseConverter<tosa::CosOp>,
3158 PointwiseConverter<tosa::TanhOp>,
3159 PointwiseConverter<tosa::ErfOp>,
3160 PointwiseConverter<tosa::BitwiseAndOp>,
3161 PointwiseConverter<tosa::BitwiseOrOp>,
3162 PointwiseConverter<tosa::BitwiseNotOp>,
3163 PointwiseConverter<tosa::BitwiseXorOp>,
3164 PointwiseConverter<tosa::LogicalAndOp>,
3165 PointwiseConverter<tosa::LogicalNotOp>,
3166 PointwiseConverter<tosa::LogicalOrOp>,
3167 PointwiseConverter<tosa::LogicalXorOp>,
3168 PointwiseConverter<tosa::CastOp>,
3169 PointwiseConverter<tosa::LogicalLeftShiftOp>,
3170 PointwiseConverter<tosa::LogicalRightShiftOp>,
3171 PointwiseConverter<tosa::ArithmeticRightShiftOp>,
3172 PointwiseConverter<tosa::ClzOp>,
3173 PointwiseConverter<tosa::SelectOp>,
3174 PointwiseConverter<tosa::GreaterOp>,
3175 PointwiseConverter<tosa::GreaterEqualOp>,
3176 PointwiseConverter<tosa::EqualOp>,
3177 PointwiseConverter<tosa::MaximumOp>,
3178 PointwiseConverter<tosa::MinimumOp>,
3179 PointwiseConverter<tosa::CeilOp>,
3180 PointwiseConverter<tosa::FloorOp>,
3181 PointwiseConverter<tosa::ClampOp>,
3182 PointwiseConverter<tosa::SigmoidOp>
3183 >(converter, patterns->getContext());
3184
3185 patterns->add<
3186 IdentityNConverter<tosa::IdentityOp>,
3187 GatherConverter,
3188 RescaleConverter,
3189 ReverseConverter,
3190 RFFT2dConverter,
3191 FFT2dConverter,
3192 TableConverter,
3193 TileConverter>(patterns->getContext());
3194
3195 // Reductions seeded with a float min/max identity need to know whether
3196 // non-finite values are available on the target.
3197 patterns->add<
3198 ReduceConverter<tosa::ReduceAllOp>,
3199 ReduceConverter<tosa::ReduceAnyOp>,
3200 ReduceConverter<tosa::ReduceMinOp>,
3201 ReduceConverter<tosa::ReduceMaxOp>,
3202 ReduceConverter<tosa::ReduceSumOp>,
3203 ReduceConverter<tosa::ReduceProductOp>,
3204 ArgMaxConverter>(patterns->getContext(), options.allowNonFinites);
3205 // clang-format on
3206}
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be inserted(the insertion happens right before the *insertion point). Since `begin` can itself be invalidated due to the memref *rewriting done from this method
static llvm::ManagedStatic< PassManagerOptions > options
static Value clamp(ImplicitLocOpBuilder &builder, Value value, Value lowerBound, Value upperBound)
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Value min(ImplicitLocOpBuilder &builder, Value value, Value bound)
static TypedAttr createInitialValueForReduceOp(Operation *op, Type elementTy, PatternRewriter &rewriter, bool allowNonFinites)
static OpFoldResult getOrFoldTensorDim(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value tensor, int64_t index)
static LogicalResult emitElementwiseComputation(ConversionPatternRewriter &rewriter, Location loc, Operation *operation, ValueRange operands, ArrayRef< OpFoldResult > targetShape, const TypeConverter &converter)
static Value createLinalgBodyCalculationForReduceOp(Operation *op, ValueRange args, Type elementTy, PatternRewriter &rewriter)
static OpTy createWithDefaultProperties(OpBuilder &builder, Location loc, TypeRange resultTypes, ValueRange operands)
static Value getTensorDim(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value tensor, int64_t index)
static Value createIndex(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, int64_t index)
static std::pair< OpFoldResult, Value > computeTargetSize(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, ValueRange operands, int64_t dim)
DenseMap< int64_t, Value > IndexPool
static LogicalResult reduceMatchAndRewriteHelper(OpTy op, uint64_t axis, PatternRewriter &rewriter, bool allowNonFinites)
static Value broadcastDynamicDimensions(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value operand, ArrayRef< OpFoldResult > targetShape, ArrayRef< Value > masterOperands)
static Value broadcastDynamicDimension(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value operand, int64_t dim, OpFoldResult targetSize, Value masterOperand)
static LogicalResult elementwiseMatchAndRewriteHelper(Operation *operation, ValueRange operands, ConversionPatternRewriter &rewriter, const TypeConverter &converter)
static APFloat getFloatMinMaxIdentity(const llvm::fltSemantics &semantics, bool negative, bool allowNonFinites)
static Value createLinalgBodyCalculationForElementwiseOp(Operation *op, ValueRange args, ArrayRef< Type > resultTypes, ConversionPatternRewriter &rewriter)
static ValueRange getBroadcastableOperands(Operation *operation, ValueRange operands)
static Value materializeBinaryNanCheckIfRequired(OpTy op, PatternRewriter &rewriter, Value lhs, Value rhs, Value result)
static std::pair< SmallVector< OpFoldResult >, SmallVector< Value > > computeTargetShape(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, ValueRange operands)
static bool operandsAndResultsRanked(Operation *operation)
static const llvm::fltSemantics * getFloatSemantics(TruncfSrcElemTypes etype)
Float semantics the element type attributes of xevm.truncf and xevm.extf stand for.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
BlockArgument getArgument(unsigned i)
Definition Block.h:154
bool getValue() const
Return the boolean value of this attribute.
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
IntegerAttr getI32IntegerAttr(int32_t value)
Definition Builders.cpp:208
FloatType getF32Type()
Definition Builders.cpp:51
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
AffineMap getMultiDimIdentityMap(unsigned rank)
Definition Builders.cpp:396
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
AffineExpr getAffineConstantExpr(int64_t constant)
Definition Builders.cpp:381
IntegerType getI64Type()
Definition Builders.cpp:73
IntegerType getI32Type()
Definition Builders.cpp:71
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
Definition Builders.h:94
BoolAttr getBoolAttr(bool value)
Definition Builders.cpp:108
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
AffineExpr getAffineDimExpr(unsigned position)
Definition Builders.cpp:373
MLIRContext * getContext() const
Definition Builders.h:56
IntegerAttr getI8IntegerAttr(int8_t value)
Definition Builders.cpp:230
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.
An attribute that represents a reference to a dense integer vector or tensor object.
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
Definition Builders.h:632
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
This class helps build Operations.
Definition Builders.h:210
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Definition Builders.h:528
This class represents a single result from folding an operation.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
unsigned getNumOperands()
Definition Operation.h:371
result_type_range getResultTypes()
Definition Operation.h:453
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
result_range getResults()
Definition Operation.h:440
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
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
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 isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getType() const
Type front()
Return first type in the range.
Definition TypeRange.h:164
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
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
Definition ArithOps.cpp:297
InFlightDiagnostic & next(InFlightDiagnostic &diag)
Starts a new message part in an in-flight diagnostic.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
OpFoldResult getMixedSize(OpBuilder &builder, Location loc, Value value, int64_t dim)
Return the dimension of the given tensor value.
Definition TensorOps.cpp:82
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Definition TensorOps.cpp:91
Value clampFloatHelper(Location loc, Value arg, Value min, Value max, OpBuilder &rewriter)
SmallVector< utils::IteratorType > getNParallelLoopsAttrs(unsigned nParallelLoops)
std::optional< SmallVector< Value > > checkHasDynamicBatchDims(PatternRewriter &rewriter, Op op, ArrayRef< Value > params)
void populateTosaToLinalgConversionPatterns(const TypeConverter &converter, RewritePatternSet *patterns, const TosaToLinalgOptions &options=TosaToLinalgOptions())
Populates conversion passes from TOSA dialect to Linalg dialect.
Value getTosaConstShape(ImplicitLocOpBuilder &builder, llvm::ArrayRef< int64_t > shape)
SmallVector< int64_t > convertFromMlirShape(ArrayRef< int64_t > shape)
Value clampIntHelper(Location loc, Value arg, Value min, Value max, OpBuilder &rewriter, bool isUnsigned)
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
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Definition Utils.cpp:114
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...