#pragma once #include "StaticIntSequence.hpp" #include "llvm/ADT/ArrayRef.h" #include "llvm/ADT/SmallVector.h" #include namespace onnx_mlir { class ConstantPool; class StaticIntGrid { public: static mlir::FailureOr fromColumns( size_t rows, llvm::ArrayRef columns, int64_t defaultValue); static mlir::FailureOr fromRows( llvm::ArrayRef rows); static mlir::FailureOr affine2D( int64_t base, int64_t rowStep, int64_t columnStep, size_t rows, size_t columns); static mlir::FailureOr laneIntervals( size_t columns, llvm::ArrayRef> 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 fromSequences( llvm::ArrayRef 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 overrideKeys; std::optional values; }; } // namespace onnx_mlir