#include "StaticIntGrid.hpp" #include "mlir/Dialect/Arith/IR/Arith.h" #include "AffineUtils.hpp" #include "ConstantUtils.hpp" #include #include using namespace mlir; namespace onnx_mlir { namespace { static std::optional cellCount(size_t rows, size_t columns) { if (!rows || !columns || rows > static_cast(std::numeric_limits::max()) / columns) return std::nullopt; return rows * columns; } static bool checkedFlatIndex( size_t row, size_t columns, size_t column, size_t &flat) { bool overflow; flat = llvm::SaturatingMultiplyAdd(row, columns, column, &overflow); return !overflow && flat <= static_cast(std::numeric_limits::max()); } static bool affineValue(int64_t base, int64_t rowStep, int64_t columnStep, size_t row, size_t column, int64_t &result) { if (row > static_cast(std::numeric_limits::max()) || column > static_cast(std::numeric_limits::max())) return false; int64_t rowValue, columnValue; return !llvm::MulOverflow(rowStep, static_cast(row), rowValue) && !llvm::MulOverflow(columnStep, static_cast(column), columnValue) && !llvm::AddOverflow(base, rowValue, result) && !llvm::AddOverflow(result, columnValue, result); } } // namespace FailureOr StaticIntGrid::fromSequences( ArrayRef input, bool columnsInput, int64_t sparseBase) { if (input.empty() || !input.front().size()) return failure(); size_t rowCount = columnsInput ? input.front().size() : input.size(); size_t columnCount = columnsInput ? input.size() : input.front().size(); auto cells = cellCount(rowCount, columnCount); if (!cells || llvm::any_of(input, [&](const StaticIntSequence &sequence) { return sequence.size() != input.front().size(); })) return failure(); StaticIntGrid result(rowCount, columnCount, input.front().valueAt(0)); if (llvm::all_equal(input)) { if (input.front().getKind() == StaticIntSequenceKind::Uniform) return result; result.kind = columnsInput ? Kind::ActionOnly : Kind::LaneOnly; result.values = input.front(); return result; } SmallVector outerBases; for (const StaticIntSequence &sequence : input) outerBases.push_back(sequence.valueAt(0)); if (llvm::all_of(input, [](const StaticIntSequence &sequence) { return sequence.getKind() == StaticIntSequenceKind::Uniform; })) { result.kind = columnsInput ? Kind::LaneOnly : Kind::ActionOnly; result.values = StaticIntSequence::fromValues(outerBases); return result; } auto innerStep = input.front().getAffineStep(); StaticIntSequence bases = StaticIntSequence::fromValues(outerBases); auto outerStep = bases.getAffineStep(); if (innerStep && outerStep && llvm::all_of(input, [&](const StaticIntSequence &sequence) { return sequence.getAffineStep() == innerStep; })) return affine2D(result.base, columnsInput ? *innerStep : *outerStep, columnsInput ? *outerStep : *innerStep, rowCount, columnCount); SmallVector values; values.reserve(*cells); for (size_t row = 0; row < rowCount; ++row) for (size_t column = 0; column < columnCount; ++column) values.push_back(columnsInput ? input[column].valueAt(row) : input[row].valueAt(column)); result.values = StaticIntSequence::fromValues(values); for (size_t index = 0; index < *cells; ++index) if (values[index] != sparseBase) result.overrideKeys.push_back(static_cast(index)); if (result.overrideKeys.size() <= *cells / 4) { result.kind = Kind::SparseLaneOverrides; result.base = sparseBase; } else { result.kind = Kind::Dense; result.overrideKeys.clear(); } return result; } FailureOr StaticIntGrid::fromRows( ArrayRef rows) { if (rows.empty() || !rows.front().size()) return failure(); return fromSequences(rows, false, rows.front().valueAt(0)); } FailureOr StaticIntGrid::fromColumns( size_t rowCount, ArrayRef columnSequences, int64_t defaultValue) { if (!cellCount(rowCount, columnSequences.size())) return failure(); SmallVector padded; padded.reserve(columnSequences.size()); for (const StaticIntSequence &sequence : columnSequences) { if (sequence.size() > rowCount) return failure(); if (sequence.size() == rowCount) { padded.push_back(sequence); continue; } SmallVector values(rowCount, defaultValue); for (size_t row = 0; row < sequence.size(); ++row) values[row] = sequence.valueAt(row); padded.push_back(StaticIntSequence::fromValues(values)); } return fromSequences(padded, true, defaultValue); } FailureOr StaticIntGrid::affine2D( int64_t base, int64_t rowStep, int64_t columnStep, size_t rows, size_t columns) { int64_t last; if (!cellCount(rows, columns) || !affineValue(base, rowStep, columnStep, rows - 1, columns - 1, last)) return failure(); StaticIntGrid result(rows, columns, base); result.kind = rowStep || columnStep ? Kind::Affine : Kind::Uniform; result.rowStep = rowStep; result.columnStep = columnStep; return result; } FailureOr StaticIntGrid::laneIntervals( size_t columns, ArrayRef> intervals, int64_t insideValue, int64_t outsideValue) { if (!columns) return failure(); SmallVector values(columns, outsideValue); for (auto [begin, end] : intervals) { if (begin > end || end > columns) return failure(); std::fill(values.begin() + begin, values.begin() + end, insideValue); } StaticIntSequence row = StaticIntSequence::fromValues(values); return fromRows(ArrayRef(row)); } int64_t StaticIntGrid::valueAt(size_t row, size_t column) const { assert(row < rows && column < columns); if (kind == Kind::Uniform) return base; if (kind == Kind::ActionOnly) return values->valueAt(row); if (kind == Kind::LaneOnly) return values->valueAt(column); if (kind == Kind::Affine) { int64_t result; bool valid = affineValue(base, rowStep, columnStep, row, column, result); assert(valid); return result; } size_t flat; bool valid = checkedFlatIndex(row, columns, column, flat); assert(valid); if (kind == Kind::SparseLaneOverrides) { auto found = llvm::lower_bound(overrideKeys, static_cast(flat)); if (found == overrideKeys.end() || *found != static_cast(flat)) return base; } assert(kind == Kind::Dense || kind == Kind::SparseLaneOverrides); return values->valueAt(flat); } Value StaticIntGrid::emitLookup(Value row, Value column, Operation *constantAnchor, ConstantPool &constants, OpBuilder &builder, Location loc) const { if (kind == Kind::Uniform) return constants.getIndex(base); if (kind == Kind::ActionOnly || kind == Kind::LaneOnly) return emitStaticIntLookup( *values, kind == Kind::ActionOnly ? row : column, constantAnchor, constants, builder, loc); Value flat = affineMulConst(builder, loc, row, columns, constantAnchor); flat = arith::AddIOp::create(builder, loc, flat, column); if (kind == Kind::Affine) { Value rowValue = affineMulConst( builder, loc, row, rowStep, constantAnchor); Value columnValue = affineMulConst( builder, loc, column, columnStep, constantAnchor); Value result = arith::AddIOp::create(builder, loc, rowValue, columnValue); return affineAddConst(builder, loc, result, base, constantAnchor); } return emitStaticIntLookup( *values, flat, constantAnchor, constants, builder, loc); } OpFoldResult StaticIntGrid::emitFoldedLookup( Value row, Value column, Operation *constantAnchor, ConstantPool &constants, OpBuilder &builder, Location loc) const { return kind == Kind::Uniform ? OpFoldResult(builder.getIndexAttr(base)) : OpFoldResult(emitLookup( row, column, constantAnchor, constants, builder, loc)); } } // namespace onnx_mlir