This commit is contained in:
@@ -0,0 +1,61 @@
|
||||
#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
|
||||
Reference in New Issue
Block a user