224 lines
8.2 KiB
C++
224 lines
8.2 KiB
C++
#include "StaticIntGrid.hpp"
|
|
|
|
#include "mlir/Dialect/Arith/IR/Arith.h"
|
|
|
|
#include "AffineUtils.hpp"
|
|
#include "ConstantUtils.hpp"
|
|
#include <algorithm>
|
|
#include <limits>
|
|
|
|
using namespace mlir;
|
|
|
|
namespace onnx_mlir {
|
|
namespace {
|
|
|
|
static std::optional<size_t> cellCount(size_t rows, size_t columns) {
|
|
if (!rows || !columns ||
|
|
rows > static_cast<size_t>(std::numeric_limits<int64_t>::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<size_t>(std::numeric_limits<int64_t>::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<size_t>(std::numeric_limits<int64_t>::max()) ||
|
|
column > static_cast<size_t>(std::numeric_limits<int64_t>::max()))
|
|
return false;
|
|
int64_t rowValue, columnValue;
|
|
return !llvm::MulOverflow(rowStep, static_cast<int64_t>(row), rowValue)
|
|
&& !llvm::MulOverflow(columnStep, static_cast<int64_t>(column),
|
|
columnValue)
|
|
&& !llvm::AddOverflow(base, rowValue, result)
|
|
&& !llvm::AddOverflow(result, columnValue, result);
|
|
}
|
|
|
|
} // namespace
|
|
|
|
FailureOr<StaticIntGrid> StaticIntGrid::fromSequences(
|
|
ArrayRef<StaticIntSequence> 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<int64_t> 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<int64_t> 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<int64_t>(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> StaticIntGrid::fromRows(
|
|
ArrayRef<StaticIntSequence> rows) {
|
|
if (rows.empty() || !rows.front().size())
|
|
return failure();
|
|
return fromSequences(rows, false, rows.front().valueAt(0));
|
|
}
|
|
|
|
FailureOr<StaticIntGrid> StaticIntGrid::fromColumns(
|
|
size_t rowCount, ArrayRef<StaticIntSequence> columnSequences,
|
|
int64_t defaultValue) {
|
|
if (!cellCount(rowCount, columnSequences.size()))
|
|
return failure();
|
|
SmallVector<StaticIntSequence> 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<int64_t> 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> 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> StaticIntGrid::laneIntervals(
|
|
size_t columns, ArrayRef<std::pair<size_t, size_t>> intervals,
|
|
int64_t insideValue, int64_t outsideValue) {
|
|
if (!columns)
|
|
return failure();
|
|
SmallVector<int64_t> 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<StaticIntSequence>(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<int64_t>(flat));
|
|
if (found == overrideKeys.end() || *found != static_cast<int64_t>(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
|