Adding buffer view encoding to the compiler. Last-mile is unplumbed (actually getting the encoding from the tensor) but the encoding is now passed from the compiler to the runtime. Progress on #6762.
diff --git a/iree/compiler/Dialect/HAL/Conversion/FlowToHAL/ConvertStreamOps.cpp b/iree/compiler/Dialect/HAL/Conversion/FlowToHAL/ConvertStreamOps.cpp index 83d6d35..9db52a9 100644 --- a/iree/compiler/Dialect/HAL/Conversion/FlowToHAL/ConvertStreamOps.cpp +++ b/iree/compiler/Dialect/HAL/Conversion/FlowToHAL/ConvertStreamOps.cpp
@@ -334,6 +334,8 @@ auto elementType = getElementType(shapedType.getElementType(), builder); assert(elementType && "unhandled element type for allocation"); + auto encodingType = getEncodingType({}, builder); + assert(encodingType && "unhandled encoding type for allocation"); SmallVector<Value> shapeDims(shapedType.getRank()); int64_t dynamicDimIndex = 0; @@ -346,7 +348,7 @@ } auto size = builder.createOrFold<IREE::HAL::AllocatorComputeSizeOp>( - loc, allocator(), shapeDims, elementType); + loc, allocator(), shapeDims, elementType, encodingType); if (shapedType.hasStaticShape()) { staticShapeToSizeMap[shapedType] = size; } else { @@ -405,6 +407,17 @@ return constantValue; } + Value getEncodingType(Attribute encodingType, OpBuilder &builder) { + auto it = memoizedEncodingTypesConstants.find(encodingType); + if (it != memoizedEncodingTypesConstants.end()) return it->second; + auto i32Value = IREE::HAL::getEncodingTypeValue(encodingType); + assert(i32Value.hasValue() && "unhandled encoding type for allocation"); + auto constantValue = + builder.createOrFold<ConstantIntOp>(loc, i32Value.getValue(), 32); + memoizedEncodingTypesConstants[encodingType] = constantValue; + return constantValue; + } + Location loc; // !hal.device used throughout the stream. @@ -429,6 +442,9 @@ // Small cache of constants used for element types. DenseMap<Type, Value> memoizedElementTypesConstants; + // Small cache of constants used for encoding types. + DenseMap<Attribute, Value> memoizedEncodingTypesConstants; + // Map of static shaped types to computed size values. DenseMap<Type, Value> staticShapeToSizeMap;
diff --git a/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertBufferViewOps.cpp b/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertBufferViewOps.cpp index 16fb3cd..32d9de2 100644 --- a/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertBufferViewOps.cpp +++ b/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertBufferViewOps.cpp
@@ -24,6 +24,8 @@ context, importSymbols, typeConverter, "hal.buffer_view.byte_length"); patterns.insert<VMImportOpConversion<IREE::HAL::BufferViewElementTypeOp>>( context, importSymbols, typeConverter, "hal.buffer_view.element_type"); + patterns.insert<VMImportOpConversion<IREE::HAL::BufferViewEncodingTypeOp>>( + context, importSymbols, typeConverter, "hal.buffer_view.encoding_type"); patterns.insert<VMImportOpConversion<IREE::HAL::BufferViewRankOp>>( context, importSymbols, typeConverter, "hal.buffer_view.rank"); patterns.insert<VMImportOpConversion<IREE::HAL::BufferViewDimOp>>(
diff --git a/iree/compiler/Dialect/HAL/Conversion/IREEToHAL/ConvertIREEToHAL.cpp b/iree/compiler/Dialect/HAL/Conversion/IREEToHAL/ConvertIREEToHAL.cpp index da50553..c50c619 100644 --- a/iree/compiler/Dialect/HAL/Conversion/IREEToHAL/ConvertIREEToHAL.cpp +++ b/iree/compiler/Dialect/HAL/Conversion/IREEToHAL/ConvertIREEToHAL.cpp
@@ -48,6 +48,11 @@ if (!elementType.hasValue()) { return rewriter.notifyMatchFailure(constantOp, "unhandled element type"); } + // TODO(#6762): get encoding type. + auto encodingType = IREE::HAL::getEncodingTypeValue({}); + if (!encodingType.hasValue()) { + return rewriter.notifyMatchFailure(constantOp, "unhandled encoding type"); + } auto buffer = rewriter.createOrFold<IREE::HAL::AllocatorConstantOp>( constantOp.getLoc(), IREE::HAL::BufferType::get(rewriter.getContext()), @@ -62,7 +67,8 @@ } auto view = rewriter.createOrFold<IREE::HAL::BufferViewCreateOp>( - constantOp.getLoc(), buffer, elementType.getValue(), shape); + constantOp.getLoc(), buffer, elementType.getValue(), + encodingType.getValue(), shape); rewriter.replaceOpWithNewOp<IREE::Util::DoNotOptimizeOp>(constantOp, view); return success();
diff --git a/iree/compiler/Dialect/HAL/Conversion/StandardToHAL/ConvertStandardToHAL.cpp b/iree/compiler/Dialect/HAL/Conversion/StandardToHAL/ConvertStandardToHAL.cpp index e04b393..2eb1ac7 100644 --- a/iree/compiler/Dialect/HAL/Conversion/StandardToHAL/ConvertStandardToHAL.cpp +++ b/iree/compiler/Dialect/HAL/Conversion/StandardToHAL/ConvertStandardToHAL.cpp
@@ -57,7 +57,7 @@ newOperands.source_dims()); newValue = rewriter.create<IREE::HAL::BufferViewCreateOp>( op.getLoc(), adaptor.getBuffer(), adaptor.getElementType(), - shapeDims); + adaptor.getEncodingType(), shapeDims); } else { newValue = adaptor.getBufferView(); }
diff --git a/iree/compiler/Dialect/HAL/IR/HALBase.td b/iree/compiler/Dialect/HAL/IR/HALBase.td index 9610e29..e6c37fb 100644 --- a/iree/compiler/Dialect/HAL/IR/HALBase.td +++ b/iree/compiler/Dialect/HAL/IR/HALBase.td
@@ -355,6 +355,10 @@ def HAL_ElementTypeAttr : SignlessIntegerAttrBase< I32, "element type attribute">; +def HAL_EncodingType : TypeAlias<I32>; +def HAL_EncodingTypeAttr : SignlessIntegerAttrBase< + I32, "encoding type attribute">; + def HAL_DeviceSize : TypeAlias<Index>; def HAL_DeviceSizeAttr : Util_IndexAttrBase<"iree_device_size_t">; def HAL_DeviceSizes : Variadic<HAL_DeviceSize>;
diff --git a/iree/compiler/Dialect/HAL/IR/HALOpFolders.cpp b/iree/compiler/Dialect/HAL/IR/HALOpFolders.cpp index 312ead6..f405b58 100644 --- a/iree/compiler/Dialect/HAL/IR/HALOpFolders.cpp +++ b/iree/compiler/Dialect/HAL/IR/HALOpFolders.cpp
@@ -60,6 +60,8 @@ // TODO(benvanik): use buffer constraints for alignment. BufferConstraintsAdaptor bufferConstraints(op.getLoc(), op.allocator()); + // TODO(#6762): switch based on op.encoding(). + auto elementSize = getElementByteCount(op.getLoc(), op.element_type(), rewriter); auto byteSize = @@ -89,6 +91,8 @@ // TODO(benvanik): use buffer constraints. BufferConstraintsAdaptor bufferConstraints(op.getLoc(), op.allocator()); + // TODO(#6762): switch based on op.encoding(). + auto offset = rewriter.createOrFold<mlir::ConstantIndexOp>(op.getLoc(), 0); for (size_t i = 0; i < op.indices().size(); ++i) { // TODO(benvanik): check error case in debug builds. @@ -146,10 +150,10 @@ auto startByteOffset = rewriter.createOrFold<AllocatorComputeOffsetOp>( op.getLoc(), rewriter.getIndexType(), op.allocator(), op.shape(), - op.element_type(), op.indices()); + op.element_type(), op.encoding_type(), op.indices()); auto endByteOffset = rewriter.createOrFold<AllocatorComputeOffsetOp>( op.getLoc(), rewriter.getIndexType(), op.allocator(), op.shape(), - op.element_type(), endIndices); + op.element_type(), op.encoding_type(), endIndices); auto elementSize = getElementByteCount(op.getLoc(), op.element_type(), rewriter); @@ -186,6 +190,11 @@ if (!elementType.hasValue()) { return rewriter.notifyMatchFailure(op, "unhandled element type"); } + // TODO(#6762): get encoding type. + auto encodingType = IREE::HAL::getEncodingTypeValue({}); + if (!encodingType.hasValue()) { + return rewriter.notifyMatchFailure(op, "unhandled encoding type"); + } // TODO(benvanik): compute from SSA use-def chain uses. IREE::HAL::MemoryTypeBitfield memoryTypes = @@ -215,7 +224,8 @@ } } auto bufferView = rewriter.createOrFold<BufferViewCreateOp>( - op.getLoc(), deviceBuffer, elementType.getValue(), shape); + op.getLoc(), deviceBuffer, elementType.getValue(), + encodingType.getValue(), shape); rewriter.replaceOp(op, {bufferView}); } else { rewriter.replaceOp(op, {deviceBuffer});
diff --git a/iree/compiler/Dialect/HAL/IR/HALOps.cpp b/iree/compiler/Dialect/HAL/IR/HALOps.cpp index 7e3fb55..edb6f90 100644 --- a/iree/compiler/Dialect/HAL/IR/HALOps.cpp +++ b/iree/compiler/Dialect/HAL/IR/HALOps.cpp
@@ -313,17 +313,19 @@ void AllocatorComputeSizeOp::build(OpBuilder &builder, OperationState &state, Value allocator, ValueRange shape, - int32_t elementType) { + int32_t elementType, int32_t encodingType) { build(builder, state, allocator, shape, - builder.createOrFold<ConstantIntOp>(state.location, elementType, 32)); + builder.createOrFold<ConstantIntOp>(state.location, elementType, 32), + builder.createOrFold<ConstantIntOp>(state.location, encodingType, 32)); } void AllocatorComputeSizeOp::build(OpBuilder &builder, OperationState &state, Value allocator, ValueRange shape, - Value elementType) { + Value elementType, Value encodingType) { state.addOperands({allocator}); state.addOperands(shape); state.addOperands(elementType); + state.addOperands(encodingType); state.addTypes({builder.getIndexType()}); } @@ -338,18 +340,22 @@ void AllocatorComputeOffsetOp::build(OpBuilder &builder, OperationState &state, Value allocator, ValueRange shape, - int32_t elementType, ValueRange indices) { + int32_t elementType, int32_t encodingType, + ValueRange indices) { build(builder, state, allocator, shape, builder.createOrFold<ConstantIntOp>(state.location, elementType, 32), + builder.createOrFold<ConstantIntOp>(state.location, encodingType, 32), indices); } void AllocatorComputeOffsetOp::build(OpBuilder &builder, OperationState &state, Value allocator, ValueRange shape, - Value elementType, ValueRange indices) { + Value elementType, Value encodingType, + ValueRange indices) { state.addOperands({allocator}); state.addOperands(shape); state.addOperands(elementType); + state.addOperands(encodingType); state.addOperands(indices); state.addTypes({builder.getIndexType()}); } @@ -365,20 +371,22 @@ void AllocatorComputeRangeOp::build(OpBuilder &builder, OperationState &state, Value allocator, ValueRange shape, - int32_t elementType, ValueRange indices, - ValueRange lengths) { + int32_t elementType, int32_t encodingType, + ValueRange indices, ValueRange lengths) { build(builder, state, allocator, shape, builder.createOrFold<ConstantIntOp>(state.location, elementType, 32), + builder.createOrFold<ConstantIntOp>(state.location, encodingType, 32), indices, lengths); } void AllocatorComputeRangeOp::build(OpBuilder &builder, OperationState &state, Value allocator, ValueRange shape, - Value elementType, ValueRange indices, - ValueRange lengths) { + Value elementType, Value encodingType, + ValueRange indices, ValueRange lengths) { state.addOperands({allocator}); state.addOperands(shape); state.addOperands(elementType); + state.addOperands(encodingType); state.addOperands(indices); state.addOperands(lengths); state.addTypes({builder.getIndexType(), builder.getIndexType()}); @@ -503,16 +511,17 @@ void BufferViewCreateOp::build(OpBuilder &builder, OperationState &state, Value buffer, int32_t elementType, - ValueRange shape) { + int32_t encodingType, ValueRange shape) { build(builder, state, buffer, builder.createOrFold<ConstantIntOp>(state.location, elementType, 32), + builder.createOrFold<ConstantIntOp>(state.location, encodingType, 32), shape); } void BufferViewCreateOp::build(OpBuilder &builder, OperationState &state, Value buffer, Value elementType, - ValueRange shape) { - state.addOperands({buffer, elementType}); + Value encodingType, ValueRange shape) { + state.addOperands({buffer, elementType, encodingType}); state.addOperands(shape); state.addTypes({BufferViewType::get(builder.getContext())}); }
diff --git a/iree/compiler/Dialect/HAL/IR/HALOps.td b/iree/compiler/Dialect/HAL/IR/HALOps.td index 91cfbf3..c5290c2 100644 --- a/iree/compiler/Dialect/HAL/IR/HALOps.td +++ b/iree/compiler/Dialect/HAL/IR/HALOps.td
@@ -117,13 +117,15 @@ let summary = [{buffer allocation size computation operation}]; let description = [{ Computes the byte size required for a buffer of the given shape and type. - This returns the same value as `hal.buffer_view.byte_length`. + The returned size may be larger than the logical size of the tensor type + stored within the buffer based on alignment requirements of the allocator. }]; let arguments = (ins HAL_Allocator:$allocator, HAL_Shape:$shape, - HAL_ElementType:$element_type + HAL_ElementType:$element_type, + HAL_EncodingType:$encoding_type ); let results = (outs HAL_DeviceSize:$result @@ -133,15 +135,24 @@ `<` $allocator `:` type($allocator) `>` `shape` `(` `[` $shape `]` `)` `type` `(` $element_type `)` + `encoding` `(` $encoding_type `)` `:` type($result) attr-dict-with-keyword }]; let builders = [ - OpBuilder<(ins "Value":$allocator, "ValueRange":$shape, - "int32_t":$elementType)>, - OpBuilder<(ins "Value":$allocator, "ValueRange":$shape, - "Value":$elementType)>, + OpBuilder<(ins + "Value":$allocator, + "ValueRange":$shape, + "int32_t":$elementType, + "int32_t":$encodingType + )>, + OpBuilder<(ins + "Value":$allocator, + "ValueRange":$shape, + "Value":$elementType, + "Value":$encodingType + )>, ]; let hasCanonicalizer = 1; @@ -160,6 +171,7 @@ HAL_Allocator:$allocator, HAL_Shape:$shape, HAL_ElementType:$element_type, + HAL_EncodingType:$encoding_type, HAL_Dims:$indices ); @@ -172,15 +184,26 @@ `indices` `(` `[` $indices `]` `)` `shape` `(` `[` $shape `]` `)` `type` `(` $element_type `)` + `encoding` `(` $encoding_type `)` `:` type($offset) attr-dict-with-keyword }]; let builders = [ - OpBuilder<(ins "Value":$allocator, "ValueRange":$shape, - "int32_t":$elementType, "ValueRange":$indices)>, - OpBuilder<(ins "Value":$allocator, "ValueRange":$shape, - "Value":$elementType, "ValueRange":$indices)>, + OpBuilder<(ins + "Value":$allocator, + "ValueRange":$shape, + "int32_t":$elementType, + "int32_t":$encodingType, + "ValueRange":$indices + )>, + OpBuilder<(ins + "Value":$allocator, + "ValueRange":$shape, + "Value":$elementType, + "Value":$encodingType, + "ValueRange":$indices + )>, ]; let hasCanonicalizer = 1; @@ -199,6 +222,7 @@ HAL_Allocator:$allocator, HAL_Shape:$shape, HAL_ElementType:$element_type, + HAL_EncodingType:$encoding_type, HAL_Dims:$indices, HAL_Dims:$lengths ); @@ -214,15 +238,28 @@ `lengths` `(` `[` $lengths `]` `)` `shape` `(` `[` $shape `]` `)` `type` `(` $element_type `)` + `encoding` `(` $encoding_type `)` `:` type($offset) `,` type($length) attr-dict-with-keyword }]; let builders = [ - OpBuilder<(ins "Value":$allocator, "ValueRange":$shape, - "int32_t":$elementType, "ValueRange":$indices, "ValueRange":$lengths)>, - OpBuilder<(ins "Value":$allocator, "ValueRange":$shape, - "Value":$elementType, "ValueRange":$indices, "ValueRange":$lengths)>, + OpBuilder<(ins + "Value":$allocator, + "ValueRange":$shape, + "int32_t":$elementType, + "int32_t":$encodingType, + "ValueRange":$indices, + "ValueRange":$lengths + )>, + OpBuilder<(ins + "Value":$allocator, + "ValueRange":$shape, + "Value":$elementType, + "Value":$encodingType, + "ValueRange":$indices, + "ValueRange":$lengths + )>, ]; let hasCanonicalizer = 1; @@ -569,6 +606,7 @@ let arguments = (ins HAL_BufferType:$buffer, HAL_ElementType:$element_type, + HAL_EncodingType:$encoding_type, HAL_Shape:$shape ); let results = (outs @@ -577,8 +615,9 @@ let assemblyFormat = [{ $buffer `,` + `shape` `=` `[` $shape `]` `,` `element_type` `=` $element_type `,` - `shape` `=` `[` $shape `]` + `encoding_type` `=` $encoding_type attr-dict `:` type($buffer) `->` type($result) }]; @@ -587,11 +626,13 @@ OpBuilder<(ins "Value":$buffer, "int32_t":$elementType, + "int32_t":$encodingType, "ValueRange":$shape )>, OpBuilder<(ins "Value":$buffer, "Value":$elementType, + "Value":$encodingType, "ValueRange":$shape )>, ]; @@ -665,6 +706,24 @@ }]; } +def HAL_BufferViewEncodingTypeOp : HAL_PureOp<"buffer_view.encoding_type"> { + let summary = [{buffer view encoding type query}]; + let description = [{ + Returns the encoding type of the buffer view. + }]; + + let arguments = (ins + HAL_BufferView:$buffer_view + ); + let results = (outs + HAL_EncodingType:$result + ); + + let assemblyFormat = [{ + $buffer_view attr-dict `:` type($result) + }]; +} + def HAL_BufferViewRankOp : HAL_PureOp<"buffer_view.rank"> { let summary = [{buffer view rank query}]; let description = [{
diff --git a/iree/compiler/Dialect/HAL/IR/HALTypes.cpp b/iree/compiler/Dialect/HAL/IR/HALTypes.cpp index a4952ec..2f9c7b6 100644 --- a/iree/compiler/Dialect/HAL/IR/HALTypes.cpp +++ b/iree/compiler/Dialect/HAL/IR/HALTypes.cpp
@@ -360,6 +360,20 @@ c8); } +llvm::Optional<int32_t> getEncodingTypeValue(Attribute attr) { + // TODO(#6762): encoding attribute handling/mapping to enums. + assert(!attr && "encoding types other than default not yet supported"); + // Default to IREE_HAL_ENCODING_TYPE_DENSE_ROW_MAJOR for now. + return 1; +} + +IntegerAttr getEncodingTypeAttr(Attribute attr, MLIRContext *context) { + auto encodingType = getEncodingTypeValue(attr); + if (!encodingType) return {}; + return IntegerAttr::get(IntegerType::get(context, 32), + encodingType.getValue()); +} + //===----------------------------------------------------------------------===// // Size-aware type utils //===----------------------------------------------------------------------===//
diff --git a/iree/compiler/Dialect/HAL/IR/HALTypes.h b/iree/compiler/Dialect/HAL/IR/HALTypes.h index d8d4608..a0abc18 100644 --- a/iree/compiler/Dialect/HAL/IR/HALTypes.h +++ b/iree/compiler/Dialect/HAL/IR/HALTypes.h
@@ -57,6 +57,13 @@ size_t getElementByteCount(IntegerAttr elementType); Value getElementByteCount(Location loc, Value elementType, OpBuilder &builder); +// Returns a stable identifier for the MLIR encoding type or 0 (opaque) if the +// type is unsupported in the ABI. +llvm::Optional<int32_t> getEncodingTypeValue(Attribute attr); + +// Returns an attribute with the MLIR encoding type or 0 (opaque). +IntegerAttr getEncodingTypeAttr(Attribute attr, MLIRContext *context); + template <typename T> inline bool allEnumBitsSet(T value, T required) { return (static_cast<uint32_t>(value) & static_cast<uint32_t>(required)) ==
diff --git a/iree/compiler/Dialect/HAL/Utils/TypeUtils.cpp b/iree/compiler/Dialect/HAL/Utils/TypeUtils.cpp index 36e091c..ebe1e99 100644 --- a/iree/compiler/Dialect/HAL/Utils/TypeUtils.cpp +++ b/iree/compiler/Dialect/HAL/Utils/TypeUtils.cpp
@@ -83,12 +83,19 @@ auto elementType = IREE::HAL::getElementTypeValue( value.getType().cast<ShapedType>().getElementType()); if (!elementType) return {}; + + // TODO(#6762): get encoding type from value. + auto encodingType = IREE::HAL::getEncodingTypeValue({}); + if (!encodingType) return {}; + auto shape = IREE::HAL::getShapeDims(loc, value, builder); if (!shape) return {}; + auto allocatorValue = builder.createOrFold<IREE::HAL::BufferAllocatorOp>( loc, IREE::HAL::AllocatorType::get(builder.getContext()), value); return builder.createOrFold<IREE::HAL::AllocatorComputeSizeOp>( - loc, allocatorValue, *shape, elementType.getValue()); + loc, allocatorValue, *shape, elementType.getValue(), + encodingType.getValue()); } // static @@ -122,6 +129,7 @@ loc, oldValue, newValue, rewriter))); return TensorRewriteAdaptor(loc, oldValue, newValue, rewriter); } + // static llvm::Optional<TensorRewriteAdaptor> TensorRewriteAdaptor::getChecked( Location loc, Value oldValue, Value newValue, @@ -162,7 +170,7 @@ auto shapeDims = getShapeDims(); if (!shapeDims) return {}; return rewriter_.createOrFold<IREE::HAL::BufferViewCreateOp>( - loc_, newValue_, getElementType(), *shapeDims); + loc_, newValue_, getElementType(), getEncodingType(), *shapeDims); } } @@ -179,6 +187,15 @@ return IREE::HAL::getElementTypeAttr(getTensorType().getElementType()); } +int32_t TensorRewriteAdaptor::getEncodingType() { + return (int32_t)getEncodingTypeAttr().getValue().getZExtValue(); +} + +IntegerAttr TensorRewriteAdaptor::getEncodingTypeAttr() { + // TODO(#6762): get encoding attribute from the tensor type. + return IREE::HAL::getEncodingTypeAttr({}, loc_.getContext()); +} + llvm::Optional<SmallVector<Value, 4>> TensorRewriteAdaptor::getShapeDims() { return IREE::HAL::getShapeDims(loc_, oldValue_, rewriter_); } @@ -189,22 +206,18 @@ } Value TensorRewriteAdaptor::getByteLength() { - if (isBufferView()) { - return rewriter_.createOrFold<IREE::HAL::BufferViewByteLengthOp>( - loc_, getBufferView()); - } else { - auto shapeDims = getShapeDims(); - if (!shapeDims) return {}; - return rewriter_.createOrFold<IREE::HAL::AllocatorComputeSizeOp>( - loc_, getAllocator(), *shapeDims, getElementType()); - } + auto shapeDims = getShapeDims(); + if (!shapeDims) return {}; + return rewriter_.createOrFold<IREE::HAL::AllocatorComputeSizeOp>( + loc_, getAllocator(), *shapeDims, getElementType(), getEncodingType()); } Value TensorRewriteAdaptor::computeOffset(ValueRange indices) { auto shapeDims = getShapeDims(); if (!shapeDims) return {}; return rewriter_.createOrFold<IREE::HAL::AllocatorComputeOffsetOp>( - loc_, getAllocator(), *shapeDims, getElementType(), indices); + loc_, getAllocator(), *shapeDims, getElementType(), getEncodingType(), + indices); } llvm::Optional<TensorRewriteAdaptor::Range> TensorRewriteAdaptor::computeRange( @@ -212,7 +225,8 @@ auto shapeDims = getShapeDims(); if (!shapeDims) return llvm::None; auto range = rewriter_.create<IREE::HAL::AllocatorComputeRangeOp>( - loc_, getAllocator(), *shapeDims, getElementType(), indices, lengths); + loc_, getAllocator(), *shapeDims, getElementType(), getEncodingType(), + indices, lengths); return Range{range.offset(), range.length()}; }
diff --git a/iree/compiler/Dialect/HAL/Utils/TypeUtils.h b/iree/compiler/Dialect/HAL/Utils/TypeUtils.h index f91304a..d352c28 100644 --- a/iree/compiler/Dialect/HAL/Utils/TypeUtils.h +++ b/iree/compiler/Dialect/HAL/Utils/TypeUtils.h
@@ -83,6 +83,10 @@ int32_t getElementType(); IntegerAttr getElementTypeAttr(); + // Returns the encoding type of the tensor as an int32 enum value. + int32_t getEncodingType(); + IntegerAttr getEncodingTypeAttr(); + // Returns the I32 shape dimensions of the tensor. llvm::Optional<SmallVector<Value, 4>> getShapeDims(); llvm::Optional<SmallVector<Value, 4>> getShapeDims(
diff --git a/iree/compiler/Dialect/HAL/hal.imports.mlir b/iree/compiler/Dialect/HAL/hal.imports.mlir index c51db3f..b30ff0a 100644 --- a/iree/compiler/Dialect/HAL/hal.imports.mlir +++ b/iree/compiler/Dialect/HAL/hal.imports.mlir
@@ -78,6 +78,7 @@ vm.import @buffer_view.create( %buffer : !vm.ref<!hal.buffer>, %element_type : i32, + %encoding_type : i32, %shape : i32 ... ) -> !vm.ref<!hal.buffer_view> attributes {nosideeffects} @@ -100,6 +101,12 @@ ) -> i32 attributes {nosideeffects} +// Returns the encoding type of the buffer view. +vm.import @buffer_view.encoding_type( + %buffer_view : !vm.ref<!hal.buffer_view>, +) -> i32 +attributes {nosideeffects} + // Returns the rank of the buffer view. vm.import @buffer_view.rank( %buffer_view : !vm.ref<!hal.buffer_view>,
diff --git a/iree/modules/hal/exports.inl b/iree/modules/hal/exports.inl index 8d8adc3..cd49cb4 100644 --- a/iree/modules/hal/exports.inl +++ b/iree/modules/hal/exports.inl
@@ -34,9 +34,10 @@ EXPORT_FN("buffer_view.buffer", iree_hal_module_buffer_view_buffer, r, r) EXPORT_FN("buffer_view.byte_length", iree_hal_module_buffer_view_byte_length, r, i) -EXPORT_FN("buffer_view.create", iree_hal_module_buffer_view_create, riCiD, r) +EXPORT_FN("buffer_view.create", iree_hal_module_buffer_view_create, riiCiD, r) EXPORT_FN("buffer_view.dim", iree_hal_module_buffer_view_dim, ri, i) EXPORT_FN("buffer_view.element_type", iree_hal_module_buffer_view_element_type, r, i) +EXPORT_FN("buffer_view.encoding_type", iree_hal_module_buffer_view_encoding_type, r, i) EXPORT_FN("buffer_view.rank", iree_hal_module_buffer_view_rank, r, i) EXPORT_FN("buffer_view.trace", iree_hal_module_buffer_view_trace, rCrD, v)
diff --git a/iree/modules/hal/module.c b/iree/modules/hal/module.c index 3e2edc0..6a74c98 100644 --- a/iree/modules/hal/module.c +++ b/iree/modules/hal/module.c
@@ -400,16 +400,14 @@ IREE_VM_ABI_EXPORT(iree_hal_module_buffer_view_create, // iree_hal_module_state_t, // - riCiD, r) { + riiCiD, r) { iree_hal_buffer_t* source_buffer = NULL; IREE_RETURN_IF_ERROR(iree_hal_buffer_check_deref(args->r0, &source_buffer)); iree_hal_element_type_t element_type = (iree_hal_element_type_t)args->i1; - // TODO(#6762): add encoding type to vm.import. - iree_hal_encoding_type_t encoding_type = - IREE_HAL_ENCODING_TYPE_DENSE_ROW_MAJOR; + iree_hal_encoding_type_t encoding_type = (iree_hal_encoding_type_t)args->i2; iree_host_size_t shape_rank = 0; iree_hal_dim_t* shape_dims = NULL; - IREE_VM_ABI_VLA_STACK_CAST(args, a2_count, a2, iree_hal_dim_t, 128, + IREE_VM_ABI_VLA_STACK_CAST(args, a3_count, a3, iree_hal_dim_t, 128, &shape_rank, &shape_dims); iree_hal_buffer_view_t* buffer_view = NULL; @@ -451,6 +449,16 @@ return iree_ok_status(); } +IREE_VM_ABI_EXPORT(iree_hal_module_buffer_view_encoding_type, // + iree_hal_module_state_t, // + r, i) { + iree_hal_buffer_view_t* buffer_view = NULL; + IREE_RETURN_IF_ERROR( + iree_hal_buffer_view_check_deref(args->r0, &buffer_view)); + rets->i0 = (uint32_t)iree_hal_buffer_view_encoding_type(buffer_view); + return iree_ok_status(); +} + IREE_VM_ABI_EXPORT(iree_hal_module_buffer_view_rank, // iree_hal_module_state_t, // r, i) {
diff --git a/iree/vm/shims.c b/iree/vm/shims.c index 8d66693..bac9544 100644 --- a/iree/vm/shims.c +++ b/iree/vm/shims.c
@@ -19,6 +19,7 @@ IREE_VM_ABI_DEFINE_SHIM(ri, r); IREE_VM_ABI_DEFINE_SHIM(ri, v); IREE_VM_ABI_DEFINE_SHIM(riCiD, r); +IREE_VM_ABI_DEFINE_SHIM(riiCiD, r); IREE_VM_ABI_DEFINE_SHIM(riCiiiD, r); IREE_VM_ABI_DEFINE_SHIM(riCrD, r); IREE_VM_ABI_DEFINE_SHIM(rii, i);
diff --git a/iree/vm/shims.h b/iree/vm/shims.h index 686ced9..d13cf8c 100644 --- a/iree/vm/shims.h +++ b/iree/vm/shims.h
@@ -296,6 +296,14 @@ iree_vm_abi_i_t a2[0]; }); +IREE_VM_ABI_VLA_STRUCT(riiCiD, a3_count, a3, { + iree_vm_ref_t r0; + int32_t i1; + int32_t i2; + iree_vm_size_t a3_count; + iree_vm_abi_i_t a3[0]; +}); + IREE_VM_ABI_VLA_STRUCT(riCrD, a2_count, a2, { iree_vm_ref_t r0; int32_t i1; @@ -388,6 +396,7 @@ IREE_VM_ABI_DECLARE_SHIM(ri, r); IREE_VM_ABI_DECLARE_SHIM(ri, v); IREE_VM_ABI_DECLARE_SHIM(riCiD, r); +IREE_VM_ABI_DECLARE_SHIM(riiCiD, r); IREE_VM_ABI_DECLARE_SHIM(riCiiiD, r); IREE_VM_ABI_DECLARE_SHIM(riCrD, r); IREE_VM_ABI_DECLARE_SHIM(rii, i);