Files
Raptor/src/PIM/Common/IR/StaticIntGrid.cpp
T
NiccoloN ab54243fda
Validate Operations / validate-operations (push) Has been cancelled
blazingly faster
2026-07-19 09:59:49 +02:00

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