25#include "llvm/ADT/STLExtras.h"
32 return memRefType.hasStaticShape() && memRefType.getLayout().isIdentity() &&
33 !llvm::is_contained(memRefType.getShape(), 0);
38struct MemRefToEmitCDialectInterface :
public ConvertToEmitCPatternInterface {
39 MemRefToEmitCDialectInterface(Dialect *dialect)
40 : ConvertToEmitCPatternInterface(dialect) {}
44 void populateConvertToEmitCConversionPatterns(
45 ConversionTarget &
target, TypeConverter &typeConverter,
46 RewritePatternSet &patterns, std::optional<bool> lowerToCpp)
const final {
54 dialect->addInterfaces<MemRefToEmitCDialectInterface>();
63struct ConvertAlloca final :
public OpConversionPattern<memref::AllocaOp> {
64 using OpConversionPattern::OpConversionPattern;
67 matchAndRewrite(memref::AllocaOp op, OpAdaptor operands,
68 ConversionPatternRewriter &rewriter)
const override {
70 if (!op.getType().hasStaticShape()) {
71 return rewriter.notifyMatchFailure(
72 op.getLoc(),
"cannot transform alloca with dynamic shape");
75 if (op.getAlignment().value_or(1) > 1) {
78 return rewriter.notifyMatchFailure(
79 op.getLoc(),
"cannot transform alloca with alignment requirement");
82 Type resultTy = getTypeConverter()->convertType(op.getType());
84 return rewriter.notifyMatchFailure(op.getLoc(),
"cannot convert type");
86 auto noInit = emitc::OpaqueAttr::get(
getContext(),
"");
88 if (op.getType().getRank() == 0) {
89 auto pointerTy = dyn_cast<emitc::PointerType>(resultTy);
90 assert(pointerTy &&
"expected rank-0 MemRef to convert to pointer");
91 Type elemTy = pointerTy.getPointee();
92 auto var = emitc::VariableOp::create(
93 rewriter, op.getLoc(), emitc::LValueType::get(elemTy), noInit);
95 auto ptr = emitc::AddressOfOp::create(rewriter, op.getLoc(), resultTy,
98 rewriter.replaceOp(op, ptr.getResult());
102 rewriter.replaceOpWithNewOp<emitc::VariableOp>(op, resultTy, noInit);
107static Value calculateMemrefTotalSizeBytes(
Location loc, MemRefType memrefType,
109 Type convertedElementType) {
111 "incompatible memref type for EmitC conversion");
113 emitc::CallOpaqueOp elementSize = emitc::CallOpaqueOp::create(
114 builder, loc, emitc::SizeTType::get(builder.
getContext()),
117 {TypeAttr::get(convertedElementType)}));
120 int64_t numElements = llvm::product_of(memrefType.getShape());
121 emitc::ConstantOp numElementsValue = emitc::ConstantOp::create(
122 builder, loc, indexType, builder.
getIndexAttr(numElements));
125 emitc::MulOp totalSizeBytes = emitc::MulOp::create(
126 builder, loc, sizeTType, elementSize.getResult(0), numElementsValue);
128 return totalSizeBytes.getResult();
131static emitc::AddressOfOp
135 emitc::ConstantOp zeroIndex = emitc::ConstantOp::create(
138 emitc::ArrayType arrayType = arrayValue.getType();
140 emitc::SubscriptOp subPtr =
142 emitc::AddressOfOp
ptr = emitc::AddressOfOp::create(
143 builder, loc, emitc::PointerType::get(arrayType.getElementType()),
150 if (isa<emitc::PointerType>(v.
getType()))
155 if (
auto cast = v.
getDefiningOp<UnrealizedConversionCastOp>())
156 if (cast.getNumOperands() == 1 &&
157 isa<emitc::PointerType>(cast.getOperand(0).getType()))
158 return cast.getOperand(0);
163 MemRefType memrefType,
172 ? emitc::ConstantOp::create(builder, idxType, builder.
getIndexAttr(0))
178 for (
auto [dim, idx] : llvm::zip(
shape.drop_front(),
indices.drop_front())) {
180 emitc::ConstantOp::create(builder, idxType, builder.
getIndexAttr(dim));
181 linearIndex = emitc::MulOp::create(builder, idxType, linearIndex, dimSize);
182 linearIndex = emitc::AddOp::create(builder, idxType, linearIndex, idx);
187struct ConvertAlloc final :
public OpConversionPattern<memref::AllocOp> {
188 using OpConversionPattern::OpConversionPattern;
190 matchAndRewrite(memref::AllocOp allocOp, OpAdaptor operands,
191 ConversionPatternRewriter &rewriter)
const override {
192 Location loc = allocOp.getLoc();
193 MemRefType memrefType = allocOp.getType();
195 return rewriter.notifyMatchFailure(
196 loc,
"incompatible memref type for EmitC conversion");
199 Type sizeTType = emitc::SizeTType::get(rewriter.getContext());
201 getTypeConverter()->convertType(memrefType.getElementType());
203 return rewriter.notifyMatchFailure(
204 loc,
"failed to convert memref element type");
206 IndexType indexType = rewriter.getIndexType();
207 Value totalSizeBytes =
208 calculateMemrefTotalSizeBytes(loc, memrefType, rewriter, elementType);
210 emitc::CallOpaqueOp allocCall;
211 StringAttr allocFunctionName;
212 Value alignmentValue;
213 SmallVector<Value, 2> argsVec;
214 if (allocOp.getAlignment()) {
216 alignmentValue = emitc::ConstantOp::create(
217 rewriter, loc, sizeTType,
218 rewriter.getIntegerAttr(indexType,
219 allocOp.getAlignment().value_or(0)));
220 argsVec.push_back(alignmentValue);
225 argsVec.push_back(totalSizeBytes);
228 allocCall = emitc::CallOpaqueOp::create(
230 emitc::PointerType::get(
231 emitc::OpaqueType::get(rewriter.getContext(),
"void")),
232 allocFunctionName, args);
234 emitc::PointerType targetPointerType = emitc::PointerType::get(elementType);
235 emitc::CastOp castOp = emitc::CastOp::create(
236 rewriter, loc, targetPointerType, allocCall.getResult(0));
238 rewriter.replaceOp(allocOp, castOp);
243struct ConvertDealloc final :
public OpConversionPattern<memref::DeallocOp> {
244 using OpConversionPattern::OpConversionPattern;
247 matchAndRewrite(memref::DeallocOp deallocOp, OpAdaptor operands,
248 ConversionPatternRewriter &rewriter)
const override {
249 Location loc = deallocOp.getLoc();
252 Value ptr = getMemRefPointer(operands.getMemref());
254 return rewriter.notifyMatchFailure(
255 loc,
"expected pointer-backed memref for EmitC deallocation");
261 Type opaqueVoidPtrType = emitc::PointerType::get(
262 emitc::OpaqueType::get(rewriter.getContext(),
"void"));
264 emitc::CastOp::create(rewriter, loc, opaqueVoidPtrType, ptr);
265 emitc::CallOpaqueOp freeCall = emitc::CallOpaqueOp::create(
268 rewriter.replaceOp(deallocOp, freeCall.getResults());
273struct ConvertCopy final :
public OpConversionPattern<memref::CopyOp> {
274 using OpConversionPattern::OpConversionPattern;
277 matchAndRewrite(memref::CopyOp copyOp, OpAdaptor operands,
278 ConversionPatternRewriter &rewriter)
const override {
279 Location loc = copyOp.getLoc();
280 MemRefType srcMemrefType = cast<MemRefType>(copyOp.getSource().getType());
281 MemRefType targetMemrefType =
282 cast<MemRefType>(copyOp.getTarget().getType());
285 return rewriter.notifyMatchFailure(
286 loc,
"incompatible source memref type for EmitC conversion");
289 return rewriter.notifyMatchFailure(
290 loc,
"incompatible target memref type for EmitC conversion");
292 if (srcMemrefType.getRank() == 0) {
293 assert(targetMemrefType.getRank() == 0 &&
294 "target must have same rank as source");
296 getTypeConverter()->convertType(srcMemrefType.getElementType());
298 return rewriter.notifyMatchFailure(loc,
"cannot convert element type");
300 Value srcPtr = getMemRefPointer(operands.getSource());
301 Value targetPtr = getMemRefPointer(operands.getTarget());
302 if (!srcPtr || !targetPtr)
303 return rewriter.notifyMatchFailure(loc,
"expected pointer operands");
305 Value zeroIndex = emitc::ConstantOp::create(
306 rewriter, loc, rewriter.getIndexType(), rewriter.getIndexAttr(0));
307 Value srcLValue = emitc::SubscriptOp::create(
311 emitc::LoadOp::create(rewriter, loc, elementType, srcLValue);
313 Value targetLValue = emitc::SubscriptOp::create(
316 rewriter.replaceOpWithNewOp<emitc::AssignOp>(copyOp, targetLValue, value);
321 cast<TypedValue<emitc::ArrayType>>(operands.getSource());
322 emitc::AddressOfOp srcPtr =
323 createPointerFromEmitcArray(loc, rewriter, srcArrayValue);
325 auto targetArrayValue =
326 cast<TypedValue<emitc::ArrayType>>(operands.getTarget());
327 emitc::AddressOfOp targetPtr =
328 createPointerFromEmitcArray(loc, rewriter, targetArrayValue);
330 Type convertedElementType =
331 getTypeConverter()->convertType(srcMemrefType.getElementType());
332 if (!convertedElementType) {
333 return rewriter.notifyMatchFailure(
334 loc,
"failed to convert memref element type");
336 Value totalSizeInBytes = calculateMemrefTotalSizeBytes(
337 loc, srcMemrefType, rewriter, convertedElementType);
338 emitc::CallOpaqueOp memCpyCall =
339 emitc::CallOpaqueOp::create(rewriter, loc,
TypeRange{},
"memcpy",
341 targetPtr.getResult(),
346 rewriter.replaceOp(copyOp, memCpyCall.getResults());
352struct ConvertGlobal final :
public OpConversionPattern<memref::GlobalOp> {
353 using OpConversionPattern::OpConversionPattern;
356 matchAndRewrite(memref::GlobalOp op, OpAdaptor operands,
357 ConversionPatternRewriter &rewriter)
const override {
358 MemRefType opTy = op.getType();
359 if (!op.getType().hasStaticShape()) {
360 return rewriter.notifyMatchFailure(
361 op.getLoc(),
"cannot transform global with dynamic shape");
364 if (op.getAlignment().value_or(1) > 1) {
366 return rewriter.notifyMatchFailure(
367 op.getLoc(),
"global variable with alignment requirement is "
368 "currently not supported");
371 Type resultTy = getTypeConverter()->convertType(opTy);
374 return rewriter.notifyMatchFailure(op.getLoc(),
375 "cannot convert result type");
379 if (visibility != SymbolTable::Visibility::Public &&
380 visibility != SymbolTable::Visibility::Private) {
381 return rewriter.notifyMatchFailure(
383 "only public and private visibility is currently supported");
387 bool staticSpecifier = visibility == SymbolTable::Visibility::Private;
388 bool externSpecifier = !staticSpecifier;
390 Attribute initialValue = operands.getInitialValueAttr();
391 if (opTy.getRank() == 0) {
392 auto pointerTy = dyn_cast<emitc::PointerType>(resultTy);
393 assert(pointerTy &&
"expected rank-0 MemRef to convert to pointer");
394 resultTy = pointerTy.getPointee();
396 if (std::optional<Attribute> initValueAttr = op.getInitialValue()) {
397 if (
auto elementsAttr = llvm::dyn_cast<ElementsAttr>(*initValueAttr)) {
398 initialValue = elementsAttr.getSplatValue<Attribute>();
402 if (isa_and_present<UnitAttr>(initialValue))
405 rewriter.replaceOpWithNewOp<emitc::GlobalOp>(
406 op, operands.getSymName(), resultTy, initialValue, externSpecifier,
407 staticSpecifier, operands.getConstant());
412struct ConvertGetGlobal final
413 :
public OpConversionPattern<memref::GetGlobalOp> {
414 using OpConversionPattern::OpConversionPattern;
417 matchAndRewrite(memref::GetGlobalOp op, OpAdaptor operands,
418 ConversionPatternRewriter &rewriter)
const override {
420 MemRefType opTy = op.getType();
421 Type resultTy = getTypeConverter()->convertType(opTy);
424 return rewriter.notifyMatchFailure(op.getLoc(),
425 "cannot convert result type");
428 if (opTy.getRank() == 0) {
429 auto pointerTy = dyn_cast<emitc::PointerType>(resultTy);
430 assert(pointerTy &&
"expected rank-0 MemRef to convert to pointer");
431 Type elemTy = pointerTy.getPointee();
432 emitc::LValueType lvalueType = emitc::LValueType::get(elemTy);
433 emitc::GetGlobalOp globalLValue = emitc::GetGlobalOp::create(
434 rewriter, op.getLoc(), lvalueType, operands.getNameAttr());
435 rewriter.replaceOpWithNewOp<emitc::AddressOfOp>(op, resultTy,
439 rewriter.replaceOpWithNewOp<emitc::GetGlobalOp>(op, resultTy,
440 operands.getNameAttr());
445struct ConvertLoad final :
public OpConversionPattern<memref::LoadOp> {
446 using OpConversionPattern::OpConversionPattern;
450 ConversionPatternRewriter &rewriter)
const override {
451 Location loc = op.getLoc();
452 auto resultTy = getTypeConverter()->convertType(op.getType());
454 return rewriter.notifyMatchFailure(loc,
"cannot convert type");
458 dyn_cast<TypedValue<emitc::ArrayType>>(operands.getMemref());
459 Value ptr = getMemRefPointer(operands.getMemref());
460 if (!ptr && arrayValue) {
461 auto subscript = emitc::SubscriptOp::create(rewriter, loc, arrayValue,
462 operands.getIndices());
464 rewriter.replaceOpWithNewOp<emitc::LoadOp>(op, resultTy, subscript);
469 return rewriter.notifyMatchFailure(loc,
"expected array or pointer type");
470 MemRefType opMemrefType = cast<MemRefType>(op.getMemref().getType());
473 ImplicitLocOpBuilder
b(loc, rewriter);
474 Value linearIndex = computeRowMajorLinearIndex(
b, opMemrefType,
indices);
475 auto typedPtr = cast<TypedValue<emitc::PointerType>>(ptr);
477 emitc::SubscriptOp::create(rewriter, loc, typedPtr, linearIndex);
479 rewriter.replaceOpWithNewOp<emitc::LoadOp>(op, resultTy, subscript);
484struct ConvertStore final :
public OpConversionPattern<memref::StoreOp> {
485 using OpConversionPattern::OpConversionPattern;
489 ConversionPatternRewriter &rewriter)
const override {
490 Location loc = op.getLoc();
492 dyn_cast<TypedValue<emitc::ArrayType>>(operands.getMemref());
493 Value ptr = getMemRefPointer(operands.getMemref());
494 if (!ptr && arrayValue) {
495 auto subscript = emitc::SubscriptOp::create(rewriter, loc, arrayValue,
496 operands.getIndices());
497 rewriter.replaceOpWithNewOp<emitc::AssignOp>(op, subscript,
498 operands.getValue());
503 return rewriter.notifyMatchFailure(loc,
"expected array or pointer type");
504 MemRefType opMemrefType = cast<MemRefType>(op.getMemref().getType());
507 ImplicitLocOpBuilder
b(loc, rewriter);
508 Value linearIndex = computeRowMajorLinearIndex(
b, opMemrefType,
indices);
509 auto typedPtr = cast<TypedValue<emitc::PointerType>>(ptr);
511 emitc::SubscriptOp::create(rewriter, loc, typedPtr, linearIndex);
513 rewriter.replaceOpWithNewOp<emitc::AssignOp>(op, subscript,
514 operands.getValue());
523 patterns.
add<ConvertAlloca, ConvertAlloc, ConvertCopy, ConvertDealloc,
static bool isMemRefTypeLegalForEmitC(MemRefType memRefType)
constexpr const char * mallocFunctionName
constexpr const char * freeFunctionName
constexpr const char * alignedAllocFunctionName
IntegerAttr getIndexAttr(int64_t value)
StringAttr getStringAttr(const Twine &bytes)
MLIRContext * getContext() const
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class helps build Operations.
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
static Visibility getSymbolVisibility(Operation *symbol)
Returns the visibility of the given symbol operation.
Visibility
An enumeration detailing the different visibility types that a symbol may have.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Include the generated interface declarations.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
void populateMemRefToEmitCConversionPatterns(RewritePatternSet &patterns, const TypeConverter &converter)
void registerConvertMemRefToEmitCInterface(DialectRegistry ®istry)
LogicalResult matchAndRewrite(spirv::LoadOp loadOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(spirv::StoreOp storeOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override