62 lines
1.8 KiB
C++
62 lines
1.8 KiB
C++
#pragma once
|
|
|
|
#include "StaticIntSequence.hpp"
|
|
|
|
#include "llvm/ADT/ArrayRef.h"
|
|
#include "llvm/ADT/SmallVector.h"
|
|
|
|
#include <utility>
|
|
|
|
namespace onnx_mlir {
|
|
|
|
class ConstantPool;
|
|
|
|
class StaticIntGrid {
|
|
public:
|
|
static mlir::FailureOr<StaticIntGrid> fromColumns(
|
|
size_t rows, llvm::ArrayRef<StaticIntSequence> columns,
|
|
int64_t defaultValue);
|
|
static mlir::FailureOr<StaticIntGrid> fromRows(
|
|
llvm::ArrayRef<StaticIntSequence> rows);
|
|
static mlir::FailureOr<StaticIntGrid> affine2D(
|
|
int64_t base, int64_t rowStep, int64_t columnStep,
|
|
size_t rows, size_t columns);
|
|
static mlir::FailureOr<StaticIntGrid> laneIntervals(
|
|
size_t columns,
|
|
llvm::ArrayRef<std::pair<size_t, size_t>> intervals,
|
|
int64_t insideValue, int64_t outsideValue);
|
|
|
|
int64_t valueAt(size_t row, size_t column) const;
|
|
|
|
mlir::Value emitLookup(mlir::Value row, mlir::Value column,
|
|
mlir::Operation *constantAnchor,
|
|
ConstantPool &constants, mlir::OpBuilder &builder,
|
|
mlir::Location loc) const;
|
|
mlir::OpFoldResult emitFoldedLookup(
|
|
mlir::Value row, mlir::Value column, mlir::Operation *constantAnchor,
|
|
ConstantPool &constants, mlir::OpBuilder &builder,
|
|
mlir::Location loc) const;
|
|
|
|
private:
|
|
enum class Kind { Uniform, ActionOnly, LaneOnly, Affine,
|
|
SparseLaneOverrides, Dense };
|
|
|
|
StaticIntGrid(size_t rows, size_t columns, int64_t base)
|
|
: rows(rows), columns(columns), base(base) {}
|
|
|
|
static mlir::FailureOr<StaticIntGrid> fromSequences(
|
|
llvm::ArrayRef<StaticIntSequence> sequences, bool columns,
|
|
int64_t sparseBase);
|
|
|
|
Kind kind = Kind::Uniform;
|
|
size_t rows = 0;
|
|
size_t columns = 0;
|
|
int64_t base = 0;
|
|
int64_t rowStep = 0;
|
|
int64_t columnStep = 0;
|
|
llvm::SmallVector<int64_t> overrideKeys;
|
|
std::optional<StaticIntSequence> values;
|
|
};
|
|
|
|
} // namespace onnx_mlir
|