blob: 60f5dca31b9ef7a8c86ec510d95234caa362851a [file] [edit]
//===- rewrite.c - Test of the rewriting C API ----------------------------===//
//
// Part of the LLVM Project, under the Apache License v2.0 with LLVM
// Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
//===----------------------------------------------------------------------===//
// RUN: mlir-capi-rewrite-test 2>&1 | FileCheck %s
#include "mlir-c/Rewrite.h"
#include "mlir-c/BuiltinTypes.h"
#include "mlir-c/IR.h"
#include <assert.h>
#include <inttypes.h>
#include <stdio.h>
MlirOperation createOperationWithName(MlirContext ctx, const char *name) {
MlirStringRef nameRef = mlirStringRefCreateFromCString(name);
MlirLocation loc = mlirLocationUnknownGet(ctx);
MlirOperationState state = mlirOperationStateGet(nameRef, loc);
MlirType indexType = mlirIndexTypeGet(ctx);
mlirOperationStateAddResults(&state, 1, &indexType);
return mlirOperationCreate(&state);
}
void testInsertionPoint(MlirContext ctx) {
// CHECK-LABEL: @testInsertionPoint
fprintf(stderr, "@testInsertionPoint\n");
const char *moduleString = "\"dialect.op1\"() : () -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
MlirOperation op1 = mlirBlockGetFirstOperation(body);
// IRRewriter create
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
// Insert before op
mlirRewriterBaseSetInsertionPointBefore(rewriter, op1);
MlirOperation op2 = createOperationWithName(ctx, "dialect.op2");
mlirRewriterBaseInsert(rewriter, op2);
// Insert after op
mlirRewriterBaseSetInsertionPointAfter(rewriter, op2);
MlirOperation op3 = createOperationWithName(ctx, "dialect.op3");
mlirRewriterBaseInsert(rewriter, op3);
MlirValue op3Res = mlirOperationGetResult(op3, 0);
// Insert after value
mlirRewriterBaseSetInsertionPointAfterValue(rewriter, op3Res);
MlirOperation op4 = createOperationWithName(ctx, "dialect.op4");
mlirRewriterBaseInsert(rewriter, op4);
// Insert at beginning of block
mlirRewriterBaseSetInsertionPointToStart(rewriter, body);
MlirOperation op5 = createOperationWithName(ctx, "dialect.op5");
mlirRewriterBaseInsert(rewriter, op5);
// Insert at end of block
mlirRewriterBaseSetInsertionPointToEnd(rewriter, body);
MlirOperation op6 = createOperationWithName(ctx, "dialect.op6");
mlirRewriterBaseInsert(rewriter, op6);
// Get insertion blocks
MlirBlock block1 = mlirRewriterBaseGetBlock(rewriter);
MlirBlock block2 = mlirRewriterBaseGetInsertionBlock(rewriter);
(void)block1;
(void)block2;
assert(body.ptr == block1.ptr);
assert(body.ptr == block2.ptr);
// clang-format off
// CHECK-NEXT: module {
// CHECK-NEXT: %{{.*}} = "dialect.op5"() : () -> index
// CHECK-NEXT: %{{.*}} = "dialect.op2"() : () -> index
// CHECK-NEXT: %{{.*}} = "dialect.op3"() : () -> index
// CHECK-NEXT: %{{.*}} = "dialect.op4"() : () -> index
// CHECK-NEXT: "dialect.op1"() : () -> ()
// CHECK-NEXT: %{{.*}} = "dialect.op6"() : () -> index
// CHECK-NEXT: }
// clang-format on
mlirOperationDump(op);
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testCreateBlock(MlirContext ctx) {
// CHECK-LABEL: @testCreateBlock
fprintf(stderr, "@testCreateBlock\n");
const char *moduleString = "\"dialect.op1\"() ({^bb0:}) : () -> ()\n"
"\"dialect.op2\"() ({^bb0:}) : () -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
MlirOperation op1 = mlirBlockGetFirstOperation(body);
MlirRegion region1 = mlirOperationGetRegion(op1, 0);
MlirBlock block1 = mlirRegionGetFirstBlock(region1);
MlirOperation op2 = mlirOperationGetNextInBlock(op1);
MlirRegion region2 = mlirOperationGetRegion(op2, 0);
MlirBlock block2 = mlirRegionGetFirstBlock(region2);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
// Create block before
MlirType indexType = mlirIndexTypeGet(ctx);
MlirLocation unknown = mlirLocationUnknownGet(ctx);
mlirRewriterBaseCreateBlockBefore(rewriter, block1, 1, &indexType, &unknown);
mlirRewriterBaseSetInsertionPointToEnd(rewriter, body);
// Clone operation
mlirRewriterBaseClone(rewriter, op1);
// Clone without regions
mlirRewriterBaseCloneWithoutRegions(rewriter, op1);
// Clone region before
mlirRewriterBaseCloneRegionBefore(rewriter, region1, block2);
mlirOperationDump(op);
// clang-format off
// CHECK-NEXT: "builtin.module"() ({
// CHECK-NEXT: "dialect.op1"() ({
// CHECK-NEXT: ^{{.*}}(%{{.*}}: index):
// CHECK-NEXT: ^{{.*}}:
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: "dialect.op2"() ({
// CHECK-NEXT: ^{{.*}}(%{{.*}}: index):
// CHECK-NEXT: ^{{.*}}:
// CHECK-NEXT: ^{{.*}}:
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: "dialect.op1"() ({
// CHECK-NEXT: ^{{.*}}(%{{.*}}: index):
// CHECK-NEXT: ^{{.*}}:
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: "dialect.op1"() ({
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: }) : () -> ()
// clang-format on
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testInlineRegionBlock(MlirContext ctx) {
// CHECK-LABEL: @testInlineRegionBlock
fprintf(stderr, "@testInlineRegionBlock\n");
const char *moduleString =
"\"dialect.op1\"() ({\n"
" ^bb0(%arg0: index):\n"
" \"dialect.op1_in1\"(%arg0) [^bb1] : (index) -> ()\n"
" ^bb1():\n"
" \"dialect.op1_in2\"() : () -> ()\n"
"}) : () -> ()\n"
"\"dialect.op2\"() ({^bb0:}) : () -> ()\n"
"\"dialect.op3\"() ({\n"
" ^bb0(%arg0: index):\n"
" \"dialect.op3_in1\"(%arg0) : (index) -> ()\n"
" ^bb1():\n"
" %x = \"dialect.op3_in2\"() : () -> index\n"
" %y = \"dialect.op3_in3\"() : () -> index\n"
"}) : () -> ()\n"
"\"dialect.op4\"() ({\n"
" ^bb0():\n"
" \"dialect.op4_in1\"() : () -> index\n"
" ^bb1(%arg0: index):\n"
" \"dialect.op4_in2\"(%arg0) : (index) -> ()\n"
"}) : () -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
MlirOperation op1 = mlirBlockGetFirstOperation(body);
MlirRegion region1 = mlirOperationGetRegion(op1, 0);
MlirOperation op2 = mlirOperationGetNextInBlock(op1);
MlirRegion region2 = mlirOperationGetRegion(op2, 0);
MlirBlock block2 = mlirRegionGetFirstBlock(region2);
MlirOperation op3 = mlirOperationGetNextInBlock(op2);
MlirRegion region3 = mlirOperationGetRegion(op3, 0);
MlirBlock block3_1 = mlirRegionGetFirstBlock(region3);
MlirBlock block3_2 = mlirBlockGetNextInRegion(block3_1);
MlirOperation op3_in2 = mlirBlockGetFirstOperation(block3_2);
MlirValue op3_in2_res = mlirOperationGetResult(op3_in2, 0);
MlirOperation op3_in3 = mlirOperationGetNextInBlock(op3_in2);
MlirOperation op4 = mlirOperationGetNextInBlock(op3);
MlirRegion region4 = mlirOperationGetRegion(op4, 0);
MlirBlock block4_1 = mlirRegionGetFirstBlock(region4);
MlirOperation op4_in1 = mlirBlockGetFirstOperation(block4_1);
MlirValue op4_in1_res = mlirOperationGetResult(op4_in1, 0);
MlirBlock block4_2 = mlirBlockGetNextInRegion(block4_1);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
// Test these three functions
mlirRewriterBaseInlineRegionBefore(rewriter, region1, block2);
mlirRewriterBaseInlineBlockBefore(rewriter, block3_1, op3_in3, 1,
&op3_in2_res);
mlirRewriterBaseMergeBlocks(rewriter, block4_2, block4_1, 1, &op4_in1_res);
mlirOperationDump(op);
// clang-format off
// CHECK-NEXT: "builtin.module"() ({
// CHECK-NEXT: "dialect.op1"() ({
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: "dialect.op2"() ({
// CHECK-NEXT: ^{{.*}}(%{{.*}}: index):
// CHECK-NEXT: "dialect.op1_in1"(%{{.*}})[^[[bb:.*]]] : (index) -> ()
// CHECK-NEXT: ^[[bb]]:
// CHECK-NEXT: "dialect.op1_in2"() : () -> ()
// CHECK-NEXT: ^{{.*}}: // no predecessors
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: "dialect.op3"() ({
// CHECK-NEXT: %{{.*}} = "dialect.op3_in2"() : () -> index
// CHECK-NEXT: "dialect.op3_in1"(%{{.*}}) : (index) -> ()
// CHECK-NEXT: %{{.*}} = "dialect.op3_in3"() : () -> index
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: "dialect.op4"() ({
// CHECK-NEXT: %{{.*}} = "dialect.op4_in1"() : () -> index
// CHECK-NEXT: "dialect.op4_in2"(%{{.*}}) : (index) -> ()
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: }) : () -> ()
// clang-format on
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testReplaceOp(MlirContext ctx) {
// CHECK-LABEL: @testReplaceOp
fprintf(stderr, "@testReplaceOp\n");
const char *moduleString =
"%x, %y, %z = \"dialect.create_values\"() : () -> (index, index, index)\n"
"%x_1, %y_1 = \"dialect.op1\"() : () -> (index, index)\n"
"\"dialect.use_op1\"(%x_1, %y_1) : (index, index) -> ()\n"
"%x_2, %y_2 = \"dialect.op2\"() : () -> (index, index)\n"
"%x_3, %y_3 = \"dialect.op3\"() : () -> (index, index)\n"
"\"dialect.use_op2\"(%x_2, %y_2) : (index, index) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
// get a handle to all operations/values
MlirOperation createValues = mlirBlockGetFirstOperation(body);
MlirValue x = mlirOperationGetResult(createValues, 0);
MlirValue z = mlirOperationGetResult(createValues, 2);
MlirOperation op1 = mlirOperationGetNextInBlock(createValues);
MlirOperation useOp1 = mlirOperationGetNextInBlock(op1);
MlirOperation op2 = mlirOperationGetNextInBlock(useOp1);
MlirOperation op3 = mlirOperationGetNextInBlock(op2);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
// Test replace op with values
MlirValue xz[2] = {x, z};
mlirRewriterBaseReplaceOpWithValues(rewriter, op1, 2, xz);
// Test replace op with op
mlirRewriterBaseReplaceOpWithOperation(rewriter, op2, op3);
mlirOperationDump(op);
// clang-format off
// CHECK-NEXT: module {
// CHECK-NEXT: %[[res:.*]]:3 = "dialect.create_values"() : () -> (index, index, index)
// CHECK-NEXT: "dialect.use_op1"(%[[res]]#0, %[[res]]#2) : (index, index) -> ()
// CHECK-NEXT: %[[res2:.*]]:2 = "dialect.op3"() : () -> (index, index)
// CHECK-NEXT: "dialect.use_op2"(%[[res2]]#0, %[[res2]]#1) : (index, index) -> ()
// CHECK-NEXT: }
// clang-format on
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testErase(MlirContext ctx) {
// CHECK-LABEL: @testErase
fprintf(stderr, "@testErase\n");
const char *moduleString = "\"dialect.op_to_erase\"() : () -> ()\n"
"\"dialect.op2\"() ({\n"
"^bb0():\n"
" \"dialect.op2_nested\"() : () -> ()"
"^block_to_erase():\n"
" \"dialect.op2_nested\"() : () -> ()"
"^bb1():\n"
" \"dialect.op2_nested\"() : () -> ()"
"}) : () -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
// get a handle to all operations/values
MlirOperation opToErase = mlirBlockGetFirstOperation(body);
MlirOperation op2 = mlirOperationGetNextInBlock(opToErase);
MlirRegion op2Region = mlirOperationGetRegion(op2, 0);
MlirBlock bb0 = mlirRegionGetFirstBlock(op2Region);
MlirBlock blockToErase = mlirBlockGetNextInRegion(bb0);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
mlirRewriterBaseEraseOp(rewriter, opToErase);
mlirRewriterBaseEraseBlock(rewriter, blockToErase);
mlirOperationDump(op);
// CHECK-NEXT: module {
// CHECK-NEXT: "dialect.op2"() ({
// CHECK-NEXT: "dialect.op2_nested"() : () -> ()
// CHECK-NEXT: ^{{.*}}:
// CHECK-NEXT: "dialect.op2_nested"() : () -> ()
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: }
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testMove(MlirContext ctx) {
// CHECK-LABEL: @testMove
fprintf(stderr, "@testMove\n");
const char *moduleString = "\"dialect.op1\"() : () -> ()\n"
"\"dialect.op2\"() ({\n"
"^bb0(%arg0: index):\n"
" \"dialect.op2_1\"(%arg0) : (index) -> ()"
"^bb1(%arg1: index):\n"
" \"dialect.op2_2\"(%arg1) : (index) -> ()"
"}) : () -> ()\n"
"\"dialect.op3\"() : () -> ()\n"
"\"dialect.op4\"() : () -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
// get a handle to all operations/values
MlirOperation op1 = mlirBlockGetFirstOperation(body);
MlirOperation op2 = mlirOperationGetNextInBlock(op1);
MlirOperation op3 = mlirOperationGetNextInBlock(op2);
MlirOperation op4 = mlirOperationGetNextInBlock(op3);
MlirRegion region2 = mlirOperationGetRegion(op2, 0);
MlirBlock block0 = mlirRegionGetFirstBlock(region2);
MlirBlock block1 = mlirBlockGetNextInRegion(block0);
// Test move operations.
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
mlirRewriterBaseMoveOpBefore(rewriter, op3, op1);
mlirRewriterBaseMoveOpAfter(rewriter, op4, op1);
mlirRewriterBaseMoveBlockBefore(rewriter, block1, block0);
mlirOperationDump(op);
// CHECK-NEXT: module {
// CHECK-NEXT: "dialect.op3"() : () -> ()
// CHECK-NEXT: "dialect.op1"() : () -> ()
// CHECK-NEXT: "dialect.op4"() : () -> ()
// CHECK-NEXT: "dialect.op2"() ({
// CHECK-NEXT: ^{{.*}}(%[[arg0:.*]]: index):
// CHECK-NEXT: "dialect.op2_2"(%[[arg0]]) : (index) -> ()
// CHECK-NEXT: ^{{.*}}(%[[arg1:.*]]: index): // no predecessors
// CHECK-NEXT: "dialect.op2_1"(%[[arg1]]) : (index) -> ()
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: }
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testOpModification(MlirContext ctx) {
// CHECK-LABEL: @testOpModification
fprintf(stderr, "@testOpModification\n");
const char *moduleString =
"%x, %y = \"dialect.op1\"() : () -> (index, index)\n"
"\"dialect.op2\"(%x) : (index) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
// get a handle to all operations/values
MlirOperation op1 = mlirBlockGetFirstOperation(body);
MlirValue y = mlirOperationGetResult(op1, 1);
MlirOperation op2 = mlirOperationGetNextInBlock(op1);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
mlirRewriterBaseStartOpModification(rewriter, op1);
mlirRewriterBaseCancelOpModification(rewriter, op1);
mlirRewriterBaseStartOpModification(rewriter, op2);
mlirOperationSetOperand(op2, 0, y);
mlirRewriterBaseFinalizeOpModification(rewriter, op2);
mlirOperationDump(op);
// CHECK-NEXT: module {
// CHECK-NEXT: %[[xy:.*]]:2 = "dialect.op1"() : () -> (index, index)
// CHECK-NEXT: "dialect.op2"(%[[xy]]#1) : (index) -> ()
// CHECK-NEXT: }
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testReplaceUses(MlirContext ctx) {
// CHECK-LABEL: @testReplaceUses
fprintf(stderr, "@testReplaceUses\n");
const char *moduleString =
// Replace values with values
"%x1, %y1, %z1 = \"dialect.op1\"() : () -> (index, index, index)\n"
"%x2, %y2, %z2 = \"dialect.op2\"() : () -> (index, index, index)\n"
"\"dialect.op1_uses\"(%x1, %y1, %z1) : (index, index, index) -> ()\n"
// Replace op with values
"%x3 = \"dialect.op3\"() : () -> index\n"
"%x4 = \"dialect.op4\"() : () -> index\n"
"\"dialect.op3_uses\"(%x3) : (index) -> ()\n"
// Replace op with op
"%x5 = \"dialect.op5\"() : () -> index\n"
"%x6 = \"dialect.op6\"() : () -> index\n"
"\"dialect.op5_uses\"(%x5) : (index) -> ()\n"
// Replace op in block;
"%x7 = \"dialect.op7\"() : () -> index\n"
"%x8 = \"dialect.op8\"() : () -> index\n"
"\"dialect.op9\"() ({\n"
"^bb0:\n"
" \"dialect.op7_uses\"(%x7) : (index) -> ()\n"
"}): () -> ()\n"
"\"dialect.op7_uses\"(%x7) : (index) -> ()\n"
// Replace value with value except in op
"%x10 = \"dialect.op10\"() : () -> index\n"
"%x11 = \"dialect.op11\"() : () -> index\n"
"\"dialect.op10_uses\"(%x10) : (index) -> ()\n"
"\"dialect.op10_uses\"(%x10) : (index) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
// get a handle to all operations/values
MlirOperation op1 = mlirBlockGetFirstOperation(body);
MlirValue x1 = mlirOperationGetResult(op1, 0);
MlirValue y1 = mlirOperationGetResult(op1, 1);
MlirValue z1 = mlirOperationGetResult(op1, 2);
MlirOperation op2 = mlirOperationGetNextInBlock(op1);
MlirValue x2 = mlirOperationGetResult(op2, 0);
MlirValue y2 = mlirOperationGetResult(op2, 1);
MlirValue z2 = mlirOperationGetResult(op2, 2);
MlirOperation op1Uses = mlirOperationGetNextInBlock(op2);
MlirOperation op3 = mlirOperationGetNextInBlock(op1Uses);
MlirOperation op4 = mlirOperationGetNextInBlock(op3);
MlirValue x4 = mlirOperationGetResult(op4, 0);
MlirOperation op3Uses = mlirOperationGetNextInBlock(op4);
MlirOperation op5 = mlirOperationGetNextInBlock(op3Uses);
MlirOperation op6 = mlirOperationGetNextInBlock(op5);
MlirOperation op5Uses = mlirOperationGetNextInBlock(op6);
MlirOperation op7 = mlirOperationGetNextInBlock(op5Uses);
MlirOperation op8 = mlirOperationGetNextInBlock(op7);
MlirValue x8 = mlirOperationGetResult(op8, 0);
MlirOperation op9 = mlirOperationGetNextInBlock(op8);
MlirRegion region9 = mlirOperationGetRegion(op9, 0);
MlirBlock block9 = mlirRegionGetFirstBlock(region9);
MlirOperation op7Uses = mlirOperationGetNextInBlock(op9);
MlirOperation op10 = mlirOperationGetNextInBlock(op7Uses);
MlirValue x10 = mlirOperationGetResult(op10, 0);
MlirOperation op11 = mlirOperationGetNextInBlock(op10);
MlirValue x11 = mlirOperationGetResult(op11, 0);
MlirOperation op10Uses1 = mlirOperationGetNextInBlock(op11);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
// Replace values
mlirRewriterBaseReplaceAllUsesWith(rewriter, x1, x2);
MlirValue y1z1[2] = {y1, z1};
MlirValue y2z2[2] = {y2, z2};
mlirRewriterBaseReplaceAllValueRangeUsesWith(rewriter, 2, y1z1, y2z2);
// Replace op with values
mlirRewriterBaseReplaceOpWithValues(rewriter, op3, 1, &x4);
// Replace op with op
mlirRewriterBaseReplaceOpWithOperation(rewriter, op5, op6);
// Replace op with op in block
mlirRewriterBaseReplaceOpUsesWithinBlock(rewriter, op7, 1, &x8, block9);
// Replace value with value except in op
mlirRewriterBaseReplaceAllUsesExcept(rewriter, x10, x11, op10Uses1);
mlirOperationDump(op);
// clang-format off
// CHECK-NEXT: module {
// CHECK-NEXT: %{{.*}}:3 = "dialect.op1"() : () -> (index, index, index)
// CHECK-NEXT: %[[res2:.*]]:3 = "dialect.op2"() : () -> (index, index, index)
// CHECK-NEXT: "dialect.op1_uses"(%[[res2]]#0, %[[res2]]#1, %[[res2]]#2) : (index, index, index) -> ()
// CHECK-NEXT: %[[res4:.*]] = "dialect.op4"() : () -> index
// CHECK-NEXT: "dialect.op3_uses"(%[[res4]]) : (index) -> ()
// CHECK-NEXT: %[[res6:.*]] = "dialect.op6"() : () -> index
// CHECK-NEXT: "dialect.op5_uses"(%[[res6]]) : (index) -> ()
// CHECK-NEXT: %[[res7:.*]] = "dialect.op7"() : () -> index
// CHECK-NEXT: %[[res8:.*]] = "dialect.op8"() : () -> index
// CHECK-NEXT: "dialect.op9"() ({
// CHECK-NEXT: "dialect.op7_uses"(%[[res8]]) : (index) -> ()
// CHECK-NEXT: }) : () -> ()
// CHECK-NEXT: "dialect.op7_uses"(%[[res7]]) : (index) -> ()
// CHECK-NEXT: %[[res10:.*]] = "dialect.op10"() : () -> index
// CHECK-NEXT: %[[res11:.*]] = "dialect.op11"() : () -> index
// CHECK-NEXT: "dialect.op10_uses"(%[[res10]]) : (index) -> ()
// CHECK-NEXT: "dialect.op10_uses"(%[[res11]]) : (index) -> ()
// CHECK-NEXT: }
// clang-format on
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
}
void testGreedyRewriteDriverConfig(MlirContext ctx) {
// CHECK-LABEL: @testGreedyRewriteDriverConfig
fprintf(stderr, "@testGreedyRewriteDriverConfig\n");
// Test config creation and destruction
MlirGreedyRewriteDriverConfig config = mlirGreedyRewriteDriverConfigCreate();
// Test all configuration setters
mlirGreedyRewriteDriverConfigSetMaxIterations(config, 5);
mlirGreedyRewriteDriverConfigSetMaxNumRewrites(config, 100);
mlirGreedyRewriteDriverConfigSetUseTopDownTraversal(config, true);
mlirGreedyRewriteDriverConfigEnableFolding(config, false);
mlirGreedyRewriteDriverConfigSetStrictness(
config, MLIR_GREEDY_REWRITE_STRICTNESS_EXISTING_OPS);
mlirGreedyRewriteDriverConfigSetRegionSimplificationLevel(
config, MLIR_GREEDY_SIMPLIFY_REGION_LEVEL_NORMAL);
mlirGreedyRewriteDriverConfigEnableConstantCSE(config, false);
// Test all configuration getters and verify values
// CHECK: MaxIterations: 5
fprintf(stderr, "MaxIterations: %" PRId64 "\n",
mlirGreedyRewriteDriverConfigGetMaxIterations(config));
// CHECK: MaxNumRewrites: 100
fprintf(stderr, "MaxNumRewrites: %" PRId64 "\n",
mlirGreedyRewriteDriverConfigGetMaxNumRewrites(config));
// CHECK: UseTopDownTraversal: 1
fprintf(stderr, "UseTopDownTraversal: %d\n",
mlirGreedyRewriteDriverConfigGetUseTopDownTraversal(config));
// CHECK: FoldingEnabled: 0
fprintf(stderr, "FoldingEnabled: %d\n",
mlirGreedyRewriteDriverConfigIsFoldingEnabled(config));
// CHECK: Strictness: 2
fprintf(stderr, "Strictness: %d\n",
mlirGreedyRewriteDriverConfigGetStrictness(config));
// CHECK: RegionSimplificationLevel: 1
fprintf(stderr, "RegionSimplificationLevel: %d\n",
mlirGreedyRewriteDriverConfigGetRegionSimplificationLevel(config));
// CHECK: ConstantCSEEnabled: 0
fprintf(stderr, "ConstantCSEEnabled: %d\n",
mlirGreedyRewriteDriverConfigIsConstantCSEEnabled(config));
// CHECK: Config test completed successfully
fprintf(stderr, "Config test completed successfully\n");
mlirGreedyRewriteDriverConfigDestroy(config);
}
void testCloneWithMapping(MlirContext ctx) {
// CHECK-LABEL: @testCloneWithMapping
fprintf(stderr, "@testCloneWithMapping\n");
const char *moduleString =
"%x, %y = \"dialect.create_values\"() : () -> (index, index)\n"
"%sum = \"dialect.add\"(%x, %y) : (index, index) -> index\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirBlock body = mlirModuleGetBody(module);
MlirOperation createValues = mlirBlockGetFirstOperation(body);
MlirValue x = mlirOperationGetResult(createValues, 0);
MlirValue y = mlirOperationGetResult(createValues, 1);
MlirOperation addOp = mlirOperationGetNextInBlock(createValues);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
mlirRewriterBaseSetInsertionPointAfter(rewriter, addOp);
// Clone addOp with a mapping that swaps x -> y, y -> x
MlirIRMapping mapping = mlirIRMappingCreate();
mlirIRMappingMapValue(mapping, x, y);
mlirIRMappingMapValue(mapping, y, x);
MlirOperation cloned =
mlirRewriterBaseCloneWithMapping(rewriter, addOp, mapping);
assert(!mlirOperationIsNull(cloned));
// Verify operands are remapped
MlirValue clonedOp0 = mlirOperationGetOperand(cloned, 0);
MlirValue clonedOp1 = mlirOperationGetOperand(cloned, 1);
assert(mlirValueEqual(clonedOp0, y));
assert(mlirValueEqual(clonedOp1, x));
mlirIRMappingDestroy(mapping);
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
// CHECK: testCloneWithMapping: PASSED
fprintf(stderr, "testCloneWithMapping: PASSED\n");
}
void testInsertionPointSaveRestore(MlirContext ctx) {
// CHECK-LABEL: @testInsertionPointSaveRestore
fprintf(stderr, "@testInsertionPointSaveRestore\n");
const char *moduleString = "\"dialect.op1\"() : () -> ()\n"
"\"dialect.op2\"() : () -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation op = mlirModuleGetOperation(module);
MlirBlock body = mlirModuleGetBody(module);
MlirOperation op1 = mlirBlockGetFirstOperation(body);
MlirOperation op2 = mlirOperationGetNextInBlock(op1);
MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
// Save an insertion point that points right before op2.
mlirRewriterBaseSetInsertionPointBefore(rewriter, op2);
MlirRewriterBaseInsertPoint saved =
mlirRewriterBaseSaveInsertionPoint(rewriter);
assert(!mlirBlockIsNull(saved.block));
assert(mlirOperationEqual(saved.operationAfter, op2));
// Move the insertion point to the end of the block. An end-of-block insertion
// point round-trips with a null `operationAfter`.
mlirRewriterBaseSetInsertionPointToEnd(rewriter, body);
MlirRewriterBaseInsertPoint endIp =
mlirRewriterBaseSaveInsertionPoint(rewriter);
assert(!mlirBlockIsNull(endIp.block));
assert(mlirOperationIsNull(endIp.operationAfter));
// Restoring the first saved insertion point makes subsequent insertions land
// before op2 again, not at the end where we just were.
mlirRewriterBaseRestoreInsertionPoint(rewriter, saved);
MlirOperation opRestored =
createOperationWithName(ctx, "dialect.op_restored");
mlirRewriterBaseInsert(rewriter, opRestored);
// Restoring the null-`operationAfter` point re-establishes end-of-block, even
// though the insertion point currently sits in the middle of the block.
mlirRewriterBaseRestoreInsertionPoint(rewriter, endIp);
assert(!mlirBlockIsNull(mlirRewriterBaseGetInsertionBlock(rewriter)));
assert(mlirOperationIsNull(
mlirRewriterBaseGetOperationAfterInsertion(rewriter)));
MlirOperation opEnd = createOperationWithName(ctx, "dialect.op_end");
mlirRewriterBaseInsert(rewriter, opEnd);
// A cleared insertion point round-trips as a null block.
mlirRewriterBaseClearInsertionPoint(rewriter);
MlirRewriterBaseInsertPoint clearedIp =
mlirRewriterBaseSaveInsertionPoint(rewriter);
assert(mlirBlockIsNull(clearedIp.block));
assert(mlirOperationIsNull(clearedIp.operationAfter));
// Restoring a cleared insertion point clears the current one.
mlirRewriterBaseSetInsertionPointToStart(rewriter, body);
assert(!mlirBlockIsNull(mlirRewriterBaseGetInsertionBlock(rewriter)));
mlirRewriterBaseRestoreInsertionPoint(rewriter, clearedIp);
assert(mlirBlockIsNull(mlirRewriterBaseGetInsertionBlock(rewriter)));
mlirOperationDump(op);
// clang-format off
// CHECK: module {
// CHECK-NEXT: "dialect.op1"() : () -> ()
// CHECK-NEXT: %{{.*}} = "dialect.op_restored"() : () -> index
// CHECK-NEXT: "dialect.op2"() : () -> ()
// CHECK-NEXT: %{{.*}} = "dialect.op_end"() : () -> index
// CHECK-NEXT: }
// clang-format on
mlirIRRewriterDestroy(rewriter);
mlirModuleDestroy(module);
// CHECK: testInsertionPointSaveRestore: PASSED
fprintf(stderr, "testInsertionPointSaveRestore: PASSED\n");
}
static MlirConversionTargetLegality dynamicLegalityAlwaysLegal(MlirOperation op,
void *userData) {
(void)op;
intptr_t *counter = (intptr_t *)userData;
++(*counter);
return MLIR_CONVERSION_TARGET_LEGALITY_LEGAL;
}
static MlirConversionTargetLegality
dynamicLegalityAlwaysIllegal(MlirOperation op, void *userData) {
(void)op;
intptr_t *counter = (intptr_t *)userData;
++(*counter);
return MLIR_CONVERSION_TARGET_LEGALITY_ILLEGAL;
}
static MlirConversionTargetLegality dynamicLegalityNoOpinion(MlirOperation op,
void *userData) {
(void)op;
intptr_t *counter = (intptr_t *)userData;
++(*counter);
return MLIR_CONVERSION_TARGET_LEGALITY_NO_OPINION;
}
// Runs a partial conversion of `moduleString` against `target` with an empty
// pattern set and returns whether it succeeded. This is what actually drives
// the registered dynamic-legality callbacks.
static bool runPartialConversion(MlirContext ctx, const char *moduleString,
MlirConversionTarget target) {
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
assert(!mlirModuleIsNull(module) && "expected module to parse");
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
mlirConversionConfigDestroy(config);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirModuleDestroy(module);
return mlirLogicalResultIsSuccess(result);
}
void testConversionTargetDynamicLegality(MlirContext ctx) {
// CHECK-LABEL: @testConversionTargetDynamicLegality
fprintf(stderr, "@testConversionTargetDynamicLegality\n");
const char *opModule = "\"dialect.op1\"() : () -> ()\n";
// addDynamicallyLegalOp: callback returning true makes the op legal, so the
// (pattern-free) partial conversion succeeds and the callback is invoked.
{
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
intptr_t counter = 0;
mlirConversionTargetAddDynamicallyLegalOp(
target, mlirStringRefCreateFromCString("dialect.op1"),
dynamicLegalityAlwaysLegal, &counter);
assert(runPartialConversion(ctx, opModule, target));
assert(counter > 0 && "legality callback must be invoked");
mlirConversionTargetDestroy(target);
}
// addDynamicallyLegalOp: callback returning false makes the op illegal. With
// no pattern to legalize it, the partial conversion fails -- proving the
// callback's return value actually drives the result.
{
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
intptr_t counter = 0;
mlirConversionTargetAddDynamicallyLegalOp(
target, mlirStringRefCreateFromCString("dialect.op1"),
dynamicLegalityAlwaysIllegal, &counter);
assert(!runPartialConversion(ctx, opModule, target));
assert(counter > 0 && "legality callback must be invoked");
mlirConversionTargetDestroy(target);
}
// addDynamicallyLegalOp composition: callbacks registered for the same op are
// chained, most-recent first. A callback returning NoOpinion abstains and
// defers to the previously-registered callback. Here the first callback marks
// the op illegal and the second abstains, so the op stays illegal (conversion
// fails) and BOTH callbacks are invoked.
{
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
intptr_t illegalCounter = 0;
intptr_t noOpinionCounter = 0;
mlirConversionTargetAddDynamicallyLegalOp(
target, mlirStringRefCreateFromCString("dialect.op1"),
dynamicLegalityAlwaysIllegal, &illegalCounter);
mlirConversionTargetAddDynamicallyLegalOp(
target, mlirStringRefCreateFromCString("dialect.op1"),
dynamicLegalityNoOpinion, &noOpinionCounter);
assert(!runPartialConversion(ctx, opModule, target));
assert(noOpinionCounter > 0 && "abstaining callback must be invoked");
assert(illegalCounter > 0 && "deferred-to callback must be invoked");
mlirConversionTargetDestroy(target);
}
// addDynamicallyLegalDialect: the callback applies to every op in the
// dialect. Returning true keeps `dialect.op1` legal -> success.
{
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
intptr_t counter = 0;
mlirConversionTargetAddDynamicallyLegalDialect(
target, mlirStringRefCreateFromCString("dialect"),
dynamicLegalityAlwaysLegal, &counter);
assert(runPartialConversion(ctx, opModule, target));
assert(counter > 0 && "dialect legality callback must be invoked");
mlirConversionTargetDestroy(target);
}
// markUnknownOpDynamicallyLegal: `dialect.op1` is unregistered and otherwise
// unmarked, so the unknown-op callback decides its legality.
{
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
intptr_t counter = 0;
mlirConversionTargetMarkUnknownOpDynamicallyLegal(
target, dynamicLegalityAlwaysLegal, &counter);
assert(runPartialConversion(ctx, opModule, target));
assert(counter > 0 && "unknown-op legality callback must be invoked");
mlirConversionTargetDestroy(target);
}
// markOpRecursivelyLegal: an op marked recursively legal short-circuits the
// walk so nested ops are never checked. Here `dialect.inner` is illegal, but
// because `dialect.outer` is recursively legal the conversion still succeeds
// and the inner op's (illegal) callback is never invoked.
{
const char *nestedModule = "\"dialect.outer\"() ({\n"
" \"dialect.inner\"() : () -> ()\n"
"}) : () -> ()\n";
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
intptr_t innerCounter = 0;
intptr_t recursiveCounter = 0;
mlirConversionTargetAddDynamicallyLegalOp(
target, mlirStringRefCreateFromCString("dialect.inner"),
dynamicLegalityAlwaysIllegal, &innerCounter);
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("dialect.outer"));
mlirConversionTargetMarkOpRecursivelyLegal(
target, mlirStringRefCreateFromCString("dialect.outer"),
dynamicLegalityAlwaysLegal, &recursiveCounter);
assert(runPartialConversion(ctx, nestedModule, target));
assert(recursiveCounter > 0 && "recursive legality callback must run");
assert(innerCounter == 0 &&
"nested op must not be visited under recursive legality");
mlirConversionTargetDestroy(target);
}
// markOpRecursivelyLegal with a NULL callback: the op is unconditionally
// recursively legal (no per-instance check), so the nested illegal op is
// still skipped and the conversion succeeds.
{
const char *nestedModule = "\"dialect.outer\"() ({\n"
" \"dialect.inner\"() : () -> ()\n"
"}) : () -> ()\n";
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
intptr_t innerCounter = 0;
mlirConversionTargetAddDynamicallyLegalOp(
target, mlirStringRefCreateFromCString("dialect.inner"),
dynamicLegalityAlwaysIllegal, &innerCounter);
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("dialect.outer"));
mlirConversionTargetMarkOpRecursivelyLegal(
target, mlirStringRefCreateFromCString("dialect.outer"), NULL, NULL);
assert(runPartialConversion(ctx, nestedModule, target));
assert(innerCounter == 0 &&
"nested op must not be visited under recursive legality");
mlirConversionTargetDestroy(target);
}
// CHECK: testConversionTargetDynamicLegality: PASSED
fprintf(stderr, "testConversionTargetDynamicLegality: PASSED\n");
}
// Type conversion callback: maps i32 -> i64 and leaves every other type
// unchanged (identity). Used by the materialization tests below.
static MlirTypeConverterConversionStatus
widenI32ToI64(MlirType type, MlirType *result, void *userData) {
(void)userData;
if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32)
*result = mlirIntegerTypeGet(mlirTypeGetContext(type), 64);
else
*result = type;
return MlirTypeConverterConversionStatusSuccess;
}
// Type conversion callback that declines i32 (returns Declined) and is the
// identity on every other type. A declining callback lets the converter fall
// through to another registered conversion function.
static MlirTypeConverterConversionStatus
declineI32(MlirType type, MlirType *result, void *userData) {
(void)userData;
if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32)
return MlirTypeConverterConversionStatusDeclined;
*result = type;
return MlirTypeConverterConversionStatusSuccess;
}
// Type conversion callback that fails on i32 (returns Failure) and is the
// identity on every other type. A failure aborts the conversion without
// trying any earlier-registered conversion function.
static MlirTypeConverterConversionStatus
failI32(MlirType type, MlirType *result, void *userData) {
(void)userData;
if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32)
return MlirTypeConverterConversionStatusFailure;
*result = type;
return MlirTypeConverterConversionStatusSuccess;
}
static MlirValue buildCastMaterialization(MlirRewriterBase rewriter,
MlirType outputType, intptr_t nInputs,
MlirValue *inputs, MlirLocation loc,
void *userData) {
intptr_t *counter = (intptr_t *)userData;
if (counter)
++(*counter);
MlirOperationState state =
mlirOperationStateGet(mlirStringRefCreateFromCString("test.cast"), loc);
mlirOperationStateAddOperands(&state, nInputs, inputs);
mlirOperationStateAddResults(&state, 1, &outputType);
MlirOperation castOp = mlirOperationCreate(&state);
mlirRewriterBaseInsert(rewriter, castOp);
return mlirOperationGetResult(castOp, 0);
}
// Source materialization callback that always declines (returns a null value)
// and records that it was consulted. Used to exercise the "this materialization
// declined, try the next one" fallback path.
static MlirValue
declineSourceMaterialization(MlirRewriterBase rewriter, MlirType outputType,
intptr_t nInputs, MlirValue *inputs,
MlirLocation loc, void *userData) {
(void)rewriter;
(void)outputType;
(void)nInputs;
(void)inputs;
(void)loc;
intptr_t *declined = (intptr_t *)userData;
if (declined)
++(*declined);
return (MlirValue){NULL};
}
// Source materialization callback that records the number of inputs it was
// invoked with (into userData) before building the `test.cast`. Used to verify
// that a 1:N replacement drives a source materialization with nInputs > 1.
static MlirValue buildSourceCastRecordInputs(MlirRewriterBase rewriter,
MlirType outputType,
intptr_t nInputs,
MlirValue *inputs,
MlirLocation loc, void *userData) {
intptr_t *observedInputs = (intptr_t *)userData;
if (observedInputs)
*observedInputs = nInputs;
MlirOperationState state =
mlirOperationStateGet(mlirStringRefCreateFromCString("test.cast"), loc);
mlirOperationStateAddOperands(&state, nInputs, inputs);
mlirOperationStateAddResults(&state, 1, &outputType);
MlirOperation castOp = mlirOperationCreate(&state);
mlirRewriterBaseInsert(rewriter, castOp);
return mlirOperationGetResult(castOp, 0);
}
// 1:1 target materialization callback. Builds a `test.cast` like
// buildCastMaterialization, but additionally inspects `originalType`: the
// counter passed as userData is only bumped when `originalType` is the expected
// original (i32) type, so a passing test proves the original type is observable
// from C.
static MlirValue buildTargetCast(MlirRewriterBase rewriter, MlirType outputType,
intptr_t nInputs, MlirValue *inputs,
MlirLocation loc, MlirType originalType,
void *userData) {
intptr_t *counter = (intptr_t *)userData;
if (counter && mlirTypeIsAInteger(originalType) &&
mlirIntegerTypeGetWidth(originalType) == 32)
++(*counter);
MlirOperationState state =
mlirOperationStateGet(mlirStringRefCreateFromCString("test.cast"), loc);
mlirOperationStateAddOperands(&state, nInputs, inputs);
mlirOperationStateAddResults(&state, 1, &outputType);
MlirOperation castOp = mlirOperationCreate(&state);
mlirRewriterBaseInsert(rewriter, castOp);
return mlirOperationGetResult(castOp, 0);
}
// 1:N type conversion callback: maps i32 -> (i16, i16) and leaves every other
// type unchanged (1:1 identity).
static MlirTypeConverterConversionStatus
splitI32(MlirType type, MlirTypeConverterConversionResults results,
void *userData) {
(void)userData;
if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32) {
MlirType i16 = mlirIntegerTypeGet(mlirTypeGetContext(type), 16);
mlirTypeConverterConversionResultsAppend(results, i16);
mlirTypeConverterConversionResultsAppend(results, i16);
} else {
mlirTypeConverterConversionResultsAppend(results, type);
}
return MlirTypeConverterConversionStatusSuccess;
}
// 1:N type conversion callback: erases i32 (converts it to zero types) and
// leaves every other type unchanged (1:1 identity).
static MlirTypeConverterConversionStatus
eraseI32(MlirType type, MlirTypeConverterConversionResults results,
void *userData) {
(void)userData;
if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32)
return MlirTypeConverterConversionStatusSuccess; // append nothing -> erase
mlirTypeConverterConversionResultsAppend(results, type);
return MlirTypeConverterConversionStatusSuccess;
}
// 1:N type conversion callback: builds one `test.cast` per requested
// output type, each consuming the (single) input, and fills `outputs`
// accordingly. Bumps the counter passed as userData.
static MlirLogicalResult buildSplitCast(MlirRewriterBase rewriter,
intptr_t nOutputTypes,
MlirType *outputTypes, intptr_t nInputs,
MlirValue *inputs, MlirLocation loc,
MlirType originalType,
MlirValue *outputs, void *userData) {
intptr_t *counter = (intptr_t *)userData;
// The original type must be the i32 operand type; it cannot be recovered from
// the (i16) output types alone, so observing it here proves it is propagated.
if (counter && mlirTypeIsAInteger(originalType) &&
mlirIntegerTypeGetWidth(originalType) == 32)
++(*counter);
for (intptr_t i = 0; i < nOutputTypes; ++i) {
MlirOperationState state =
mlirOperationStateGet(mlirStringRefCreateFromCString("test.cast"), loc);
mlirOperationStateAddOperands(&state, nInputs, inputs);
mlirOperationStateAddResults(&state, 1, &outputTypes[i]);
MlirOperation castOp = mlirOperationCreate(&state);
mlirRewriterBaseInsert(rewriter, castOp);
outputs[i] = mlirOperationGetResult(castOp, 0);
}
return mlirLogicalResultSuccess();
}
// 1:N type conversion callback that appends bogus result types and then
// declines (returns Declined). The binding must roll back the appended types so
// they do not pollute a subsequently-tried conversion function.
static MlirTypeConverterConversionStatus appendThenDeclineConversion(
MlirType type, MlirTypeConverterConversionResults results, void *userData) {
intptr_t *counter = (intptr_t *)userData;
if (counter)
++(*counter);
MlirType i8 = mlirIntegerTypeGet(mlirTypeGetContext(type), 8);
// Append more (and differently-typed) entries than the real conversion would,
// so any leak is observable as a wrong type/arity downstream.
mlirTypeConverterConversionResultsAppend(results, i8);
mlirTypeConverterConversionResultsAppend(results, i8);
mlirTypeConverterConversionResultsAppend(results, i8);
return MlirTypeConverterConversionStatusDeclined;
}
// 1:N type conversion callback that appends bogus result types and then
// fails (returns Failure). Unlike a decline, this must abort the conversion
// without trying any earlier-registered conversion function; the binding must
// still roll back the appended types.
static MlirTypeConverterConversionStatus appendThenFailConversion(
MlirType type, MlirTypeConverterConversionResults results, void *userData) {
intptr_t *counter = (intptr_t *)userData;
if (counter)
++(*counter);
MlirType i8 = mlirIntegerTypeGet(mlirTypeGetContext(type), 8);
mlirTypeConverterConversionResultsAppend(results, i8);
mlirTypeConverterConversionResultsAppend(results, i8);
mlirTypeConverterConversionResultsAppend(results, i8);
return MlirTypeConverterConversionStatusFailure;
}
// 1:N target materialization that always declines (returns failure) and records
// that it was consulted. Exercises the "decline, try the next" fallback.
static MlirLogicalResult declineTargetMaterialization(
MlirRewriterBase rewriter, intptr_t nOutputTypes, MlirType *outputTypes,
intptr_t nInputs, MlirValue *inputs, MlirLocation loc,
MlirType originalType, MlirValue *outputs, void *userData) {
(void)rewriter;
(void)nOutputTypes;
(void)outputTypes;
(void)nInputs;
(void)inputs;
(void)loc;
(void)originalType;
(void)outputs;
intptr_t *counter = (intptr_t *)userData;
if (counter)
++(*counter);
return mlirLogicalResultFailure();
}
// Conversion pattern for `test.source`: replaces it with a `test.source_i64`
// op whose result has the widened (i64) type. Because the original result type
// (i32) differs from the replacement type (i64), persisting uses force the
// framework to insert a source materialization.
static MlirLogicalResult convertSource(MlirConversionPattern pattern,
MlirOperation op, intptr_t nOperands,
MlirValue *operands,
MlirConversionPatternRewriter rewriter,
void *userData) {
(void)pattern;
(void)nOperands;
(void)operands;
(void)userData;
MlirContext ctx = mlirOperationGetContext(op);
MlirLocation loc = mlirOperationGetLocation(op);
MlirType i64 = mlirIntegerTypeGet(ctx, 64);
MlirOperationState state = mlirOperationStateGet(
mlirStringRefCreateFromCString("test.source_i64"), loc);
mlirOperationStateAddResults(&state, 1, &i64);
MlirOperation newOp = mlirOperationCreate(&state);
MlirRewriterBase base = mlirPatternRewriterAsBase(
mlirConversionPatternRewriterAsPatternRewriter(rewriter));
mlirRewriterBaseInsert(base, newOp);
MlirValue newVal = mlirOperationGetResult(newOp, 0);
mlirRewriterBaseReplaceOpWithValues(base, op, 1, &newVal);
return mlirLogicalResultSuccess();
}
void testTypeConverterSourceMaterialization(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverterSourceMaterialization
fprintf(stderr, "@testTypeConverterSourceMaterialization\n");
// `test.source` produces an i32 that is consumed by the (legal) `test.user`.
// Converting `test.source` to an i64-producing op leaves `test.user` wanting
// the original i32, which triggers a source materialization back to i32.
const char *moduleString = "%0 = \"test.source\"() : () -> i32\n"
"\"test.user\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAddConversion(converter, widenI32ToI64, NULL);
intptr_t materializationCounter = 0;
mlirTypeConverterAddSourceMaterialization(converter, buildCastMaterialization,
&materializationCounter);
// Register a second materialization that always declines. Because
// materializations are tried most-recently-added first, this one runs first,
// returns null, and the framework must fall back to buildCastMaterialization.
intptr_t declinedCounter = 0;
mlirTypeConverterAddSourceMaterialization(
converter, declineSourceMaterialization, &declinedCounter);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertSource, NULL};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.source"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.source"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.source_i64"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.cast"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.user"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
assert(materializationCounter > 0 &&
"source materialization callback must be invoked");
assert(declinedCounter > 0 &&
"declining materialization must be consulted before the fallback");
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: module {
// CHECK-NEXT: %[[v:.*]] = "test.source_i64"() : () -> i64
// CHECK-NEXT: %[[c:.*]] = "test.cast"(%[[v]]) : (i64) -> i32
// CHECK-NEXT: "test.user"(%[[c]]) : (i32) -> ()
// CHECK-NEXT: }
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverterSourceMaterialization: PASSED
fprintf(stderr, "testTypeConverterSourceMaterialization: PASSED\n");
}
// Conversion pattern for `test.consumer`: replaces it with a
// `test.consumer_legal` op that consumes the (already remapped) operands. The
// operand of the original op has type i32 but its producer is not converted, so
// the framework inserts a target materialization to i64 before invoking this
// pattern -- the remapped `operands` are therefore the i64 cast results.
static MlirLogicalResult convertConsumer(MlirConversionPattern pattern,
MlirOperation op, intptr_t nOperands,
MlirValue *operands,
MlirConversionPatternRewriter rewriter,
void *userData) {
(void)pattern;
(void)userData;
MlirLocation loc = mlirOperationGetLocation(op);
MlirOperationState state = mlirOperationStateGet(
mlirStringRefCreateFromCString("test.consumer_legal"), loc);
mlirOperationStateAddOperands(&state, nOperands, operands);
MlirOperation newOp = mlirOperationCreate(&state);
MlirRewriterBase base = mlirPatternRewriterAsBase(
mlirConversionPatternRewriterAsPatternRewriter(rewriter));
mlirRewriterBaseInsert(base, newOp);
mlirRewriterBaseEraseOp(base, op);
return mlirLogicalResultSuccess();
}
// 1:N conversion pattern for `test.consumer`: builds `test.consumer_legal` from
// the flattened remapped operands. Under the i32 -> (i16, i16) conversion the
// single operand is remapped to two values, so this is invoked via the 1:N
// matchAndRewrite hook with nOperands == 2.
static MlirLogicalResult
convertConsumer1ToN(MlirConversionPattern pattern, MlirOperation op,
intptr_t nRanges, intptr_t *rangeSizes, intptr_t nOperands,
MlirValue *operands, MlirConversionPatternRewriter rewriter,
void *userData) {
(void)pattern;
(void)nRanges;
(void)rangeSizes;
(void)userData;
MlirLocation loc = mlirOperationGetLocation(op);
MlirOperationState state = mlirOperationStateGet(
mlirStringRefCreateFromCString("test.consumer_legal"), loc);
// Under an erasure conversion (i32 -> zero types) the operand range is empty;
// avoid passing a null pointer with count 0 (UB caught by UBSan's memcpy
// nonnull check).
if (nOperands)
mlirOperationStateAddOperands(&state, nOperands, operands);
MlirOperation newOp = mlirOperationCreate(&state);
MlirRewriterBase base = mlirPatternRewriterAsBase(
mlirConversionPatternRewriterAsPatternRewriter(rewriter));
mlirRewriterBaseInsert(base, newOp);
mlirRewriterBaseEraseOp(base, op);
return mlirLogicalResultSuccess();
}
void testTypeConverterTargetMaterialization(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverterTargetMaterialization
fprintf(stderr, "@testTypeConverterTargetMaterialization\n");
// `test.consumer` takes an i32 from the (legal, unconverted) `test.producer`.
// Converting `test.consumer` requires its operand as i64, which triggers a
// target materialization from i32 to i64.
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.consumer\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAddConversion(converter, widenI32ToI64, NULL);
intptr_t materializationCounter = 0;
mlirTypeConverterAddTargetMaterialization(converter, buildTargetCast,
&materializationCounter);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertConsumer,
NULL};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.consumer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.consumer_legal"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.cast"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
assert(materializationCounter > 0 &&
"target materialization callback must be invoked");
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: module {
// CHECK-NEXT: %[[v:.*]] = "test.producer"() : () -> i32
// CHECK-NEXT: %[[c:.*]] = "test.cast"(%[[v]]) : (i32) -> i64
// CHECK-NEXT: "test.consumer_legal"(%[[c]]) : (i64) -> ()
// CHECK-NEXT: }
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverterTargetMaterialization: PASSED
fprintf(stderr, "testTypeConverterTargetMaterialization: PASSED\n");
}
// Unit test for the three MlirTypeConverterConversionStatus values, exercised
// directly through mlirTypeConverterConvertType (no conversion driver needed).
// Conversion functions are consulted most-recently-registered first, so the
// second-registered callback is tried before the first (`widenI32ToI64`).
void testTypeConverterConversionStatus(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverterConversionStatus
fprintf(stderr, "@testTypeConverterConversionStatus\n");
MlirType i32 = mlirIntegerTypeGet(ctx, 32);
MlirType i64 = mlirIntegerTypeGet(ctx, 64);
// Success: the sole conversion widens i32 -> i64.
MlirTypeConverter success = mlirTypeConverterCreate();
mlirTypeConverterAddConversion(success, widenI32ToI64, NULL);
MlirType successResult = mlirTypeConverterConvertType(success, i32);
assert(!mlirTypeIsNull(successResult) && mlirTypeEqual(successResult, i64) &&
"Success must yield the converted type");
mlirTypeConverterDestroy(success);
// Declined: `declineI32` is tried first and declines, so the converter falls
// through to `widenI32ToI64`, which converts i32 -> i64.
MlirTypeConverter declined = mlirTypeConverterCreate();
mlirTypeConverterAddConversion(declined, widenI32ToI64, NULL);
mlirTypeConverterAddConversion(declined, declineI32, NULL);
MlirType declinedResult = mlirTypeConverterConvertType(declined, i32);
assert(!mlirTypeIsNull(declinedResult) &&
mlirTypeEqual(declinedResult, i64) &&
"Declined must fall through to the next conversion function");
mlirTypeConverterDestroy(declined);
// Failure: `failI32` is tried first and fails, which aborts the
// conversion without consulting `widenI32ToI64`; convertType returns null.
MlirTypeConverter failure = mlirTypeConverterCreate();
mlirTypeConverterAddConversion(failure, widenI32ToI64, NULL);
mlirTypeConverterAddConversion(failure, failI32, NULL);
MlirType failureResult = mlirTypeConverterConvertType(failure, i32);
assert(mlirTypeIsNull(failureResult) &&
"Failure must abort without trying another conversion function");
mlirTypeConverterDestroy(failure);
// CHECK: testTypeConverterConversionStatus: PASSED
fprintf(stderr, "testTypeConverterConversionStatus: PASSED\n");
}
void testTypeConverter1ToNTargetMaterialization(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverter1ToNTargetMaterialization
fprintf(stderr, "@testTypeConverter1ToNTargetMaterialization\n");
// `test.consumer` takes an i32 from the (legal, unconverted) `test.producer`.
// The 1:N conversion maps i32 -> (i16, i16), so converting `test.consumer`
// requires its operand as two i16 values, which triggers a 1:N target
// materialization producing two values from the single i32.
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.consumer\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAdd1ToNConversion(converter, splitI32, NULL);
intptr_t materializationCounter = 0;
mlirTypeConverterAdd1ToNTargetMaterialization(converter, buildSplitCast,
&materializationCounter);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, NULL,
convertConsumer1ToN};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.consumer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.consumer_legal"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.cast"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
assert(materializationCounter > 0 &&
"1:N target materialization callback must be invoked");
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: %[[v:.*]] = "test.producer"() : () -> i32
// CHECK: %[[c0:.*]] = "test.cast"(%[[v]]) : (i32) -> i16
// CHECK: %[[c1:.*]] = "test.cast"(%[[v]]) : (i32) -> i16
// CHECK: "test.consumer_legal"(%[[c0]], %[[c1]]) : (i16, i16) -> ()
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverter1ToNTargetMaterialization: PASSED
fprintf(stderr, "testTypeConverter1ToNTargetMaterialization: PASSED\n");
}
void testTypeConverter1ToNConversionDeclineRollback(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverter1ToNConversionDeclineRollback
fprintf(stderr, "@testTypeConverter1ToNConversionDeclineRollback\n");
// Same setup as the 1:N target materialization test, but with an extra
// conversion that appends bogus types and then declines. It is registered
// last, so it is tried first; the binding must roll back its appended types
// so the i32 -> (i16, i16) conversion still produces exactly two i16 values
// (not i8s, and not three of them).
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.consumer\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAdd1ToNConversion(converter, splitI32, NULL);
intptr_t declineCounter = 0;
mlirTypeConverterAdd1ToNConversion(converter, appendThenDeclineConversion,
&declineCounter);
intptr_t materializationCounter = 0;
mlirTypeConverterAdd1ToNTargetMaterialization(converter, buildSplitCast,
&materializationCounter);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, NULL,
convertConsumer1ToN};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.consumer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.consumer_legal"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.cast"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
assert(declineCounter > 0 && "declining conversion must be consulted");
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: %[[v:.*]] = "test.producer"() : () -> i32
// CHECK: %[[c0:.*]] = "test.cast"(%[[v]]) : (i32) -> i16
// CHECK: %[[c1:.*]] = "test.cast"(%[[v]]) : (i32) -> i16
// CHECK: "test.consumer_legal"(%[[c0]], %[[c1]]) : (i16, i16) -> ()
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverter1ToNConversionDeclineRollback: PASSED
fprintf(stderr, "testTypeConverter1ToNConversionDeclineRollback: PASSED\n");
}
void testTypeConverter1ToNConversionFailure(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverter1ToNConversionFailure
fprintf(stderr, "@testTypeConverter1ToNConversionFailure\n");
// Same setup as the decline-rollback test, but the extra conversion returns a
// failure instead of a decline. Because it is registered last (tried
// first), the failure must abort the conversion of i32 without falling back
// to the i32 -> (i16, i16) conversion, so the whole conversion fails and the
// IR is left unchanged.
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.consumer\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAdd1ToNConversion(converter, splitI32, NULL);
intptr_t failCounter = 0;
mlirTypeConverterAdd1ToNConversion(converter, appendThenFailConversion,
&failCounter);
mlirTypeConverterAdd1ToNTargetMaterialization(converter, buildSplitCast,
NULL);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, NULL,
convertConsumer1ToN};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.consumer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.consumer_legal"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.cast"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsFailure(result) &&
"failure must abort the conversion");
assert(failCounter > 0 && "failing conversion must be consulted");
// The conversion failed, so the IR is rolled back and left unchanged: i32 is
// never split and no `test.cast` is produced.
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: %[[v:.*]] = "test.producer"() : () -> i32
// CHECK: "test.consumer"(%[[v]]) : (i32) -> ()
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverter1ToNConversionFailure: PASSED
fprintf(stderr, "testTypeConverter1ToNConversionFailure: PASSED\n");
}
void testTypeConverter1ToNTargetMaterializationDecline(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverter1ToNTargetMaterializationDecline
fprintf(stderr, "@testTypeConverter1ToNTargetMaterializationDecline\n");
// Register two 1:N target materializations. Tried most-recently-added
// first: one that returns failure (declining), and finally the real one.
// The conversion must still succeed via the last, and both must run.
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.consumer\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAdd1ToNConversion(converter, splitI32, NULL);
intptr_t buildCounter = 0, declineCounter = 0;
mlirTypeConverterAdd1ToNTargetMaterialization(converter, buildSplitCast,
&buildCounter);
mlirTypeConverterAdd1ToNTargetMaterialization(
converter, declineTargetMaterialization, &declineCounter);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, NULL,
convertConsumer1ToN};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.consumer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.consumer_legal"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.cast"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
assert(declineCounter > 0 && "failing materialization must be consulted");
assert(buildCounter > 0 && "fallback materialization must run");
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: %[[v:.*]] = "test.producer"() : () -> i32
// CHECK: %[[c0:.*]] = "test.cast"(%[[v]]) : (i32) -> i16
// CHECK: %[[c1:.*]] = "test.cast"(%[[v]]) : (i32) -> i16
// CHECK: "test.consumer_legal"(%[[c0]], %[[c1]]) : (i16, i16) -> ()
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverter1ToNTargetMaterializationDecline: PASSED
fprintf(stderr,
"testTypeConverter1ToNTargetMaterializationDecline: PASSED\n");
}
// Conversion pattern for `test.producer`: replaces its single i32 result with
// two i16 values (built by two `test.piece` ops) via replaceOpWithMultiple.
// Because a legal use (`test.user`) still wants the original i32, the framework
// must insert a source materialization with two inputs to reconcile the 1:N
// replacement.
static MlirLogicalResult
convertProducerToMultiple(MlirConversionPattern pattern, MlirOperation op,
intptr_t nOperands, MlirValue *operands,
MlirConversionPatternRewriter rewriter,
void *userData) {
(void)pattern;
(void)nOperands;
(void)operands;
(void)userData;
MlirContext ctx = mlirOperationGetContext(op);
MlirLocation loc = mlirOperationGetLocation(op);
MlirType i16 = mlirIntegerTypeGet(ctx, 16);
MlirRewriterBase base = mlirPatternRewriterAsBase(
mlirConversionPatternRewriterAsPatternRewriter(rewriter));
MlirValue pieces[2];
for (int i = 0; i < 2; ++i) {
MlirOperationState state = mlirOperationStateGet(
mlirStringRefCreateFromCString("test.piece"), loc);
mlirOperationStateAddResults(&state, 1, &i16);
MlirOperation pieceOp = mlirOperationCreate(&state);
mlirRewriterBaseInsert(base, pieceOp);
pieces[i] = mlirOperationGetResult(pieceOp, 0);
}
// The op has a single result, so there is one range, carrying two values.
intptr_t rangeSizes[1] = {2};
mlirConversionPatternRewriterReplaceOpWithMultiple(rewriter, op, 1,
rangeSizes, pieces);
return mlirLogicalResultSuccess();
}
void testTypeConverterMultiInputSourceMaterialization(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverterMultiInputSourceMaterialization
fprintf(stderr, "@testTypeConverterMultiInputSourceMaterialization\n");
// `test.producer` produces an i32 consumed by the (legal) `test.user`.
// The pattern replaces the producer's result with two i16 values, so the
// i32-wanting `test.user` forces a source materialization that receives both
// values as inputs (nInputs == 2).
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.user\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
intptr_t observedInputs = 0;
mlirTypeConverterAddSourceMaterialization(
converter, buildSourceCastRecordInputs, &observedInputs);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL,
convertProducerToMultiple, NULL};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.producer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.piece"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.cast"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.user"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
assert(observedInputs == 2 &&
"source materialization must receive both replacement values");
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: %[[a:.*]] = "test.piece"() : () -> i16
// CHECK: %[[b:.*]] = "test.piece"() : () -> i16
// CHECK: %[[c:.*]] = "test.cast"(%[[a]], %[[b]]) : (i16, i16) -> i32
// CHECK: "test.user"(%[[c]]) : (i32) -> ()
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverterMultiInputSourceMaterialization: PASSED
fprintf(stderr, "testTypeConverterMultiInputSourceMaterialization: PASSED\n");
}
// Conversion pattern for a two-result `test.multi`: replaces it via
// replaceOpWithMultiple with two ranges (one per result), exercising the
// nRanges > 1 marshalling path. Each result is replaced 1:1 with a same-typed
// value so no materialization is needed.
static MlirLogicalResult
convertMultiResult(MlirConversionPattern pattern, MlirOperation op,
intptr_t nOperands, MlirValue *operands,
MlirConversionPatternRewriter rewriter, void *userData) {
(void)pattern;
(void)nOperands;
(void)operands;
(void)userData;
MlirContext ctx = mlirOperationGetContext(op);
MlirLocation loc = mlirOperationGetLocation(op);
MlirType i32 = mlirIntegerTypeGet(ctx, 32);
MlirType i64 = mlirIntegerTypeGet(ctx, 64);
MlirRewriterBase base = mlirPatternRewriterAsBase(
mlirConversionPatternRewriterAsPatternRewriter(rewriter));
MlirValue vals[2];
MlirType types[2] = {i32, i64};
const char *names[2] = {"test.r0", "test.r1"};
for (int i = 0; i < 2; ++i) {
MlirOperationState state =
mlirOperationStateGet(mlirStringRefCreateFromCString(names[i]), loc);
mlirOperationStateAddResults(&state, 1, &types[i]);
MlirOperation newOp = mlirOperationCreate(&state);
mlirRewriterBaseInsert(base, newOp);
vals[i] = mlirOperationGetResult(newOp, 0);
}
// Two ranges (one per result), each carrying a single value.
intptr_t rangeSizes[2] = {1, 1};
mlirConversionPatternRewriterReplaceOpWithMultiple(rewriter, op, 2,
rangeSizes, vals);
return mlirLogicalResultSuccess();
}
void testConversionReplaceOpWithMultipleRanges(MlirContext ctx) {
// CHECK-LABEL: @testConversionReplaceOpWithMultipleRanges
fprintf(stderr, "@testConversionReplaceOpWithMultipleRanges\n");
const char *moduleString = "%0:2 = \"test.multi\"() : () -> (i32, i64)\n"
"\"test.use\"(%0#0, %0#1) : (i32, i64) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertMultiResult,
NULL};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.multi"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.multi"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.r0"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.r1"));
mlirConversionTargetAddLegalOp(target,
mlirStringRefCreateFromCString("test.use"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: %[[a:.*]] = "test.r0"() : () -> i32
// CHECK: %[[b:.*]] = "test.r1"() : () -> i64
// CHECK: "test.use"(%[[a]], %[[b]]) : (i32, i64) -> ()
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testConversionReplaceOpWithMultipleRanges: PASSED
fprintf(stderr, "testConversionReplaceOpWithMultipleRanges: PASSED\n");
}
void testTypeConverter1ToNOperandRequires1ToNCallback(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverter1ToNOperandRequires1ToNCallback
fprintf(stderr, "@testTypeConverter1ToNOperandRequires1ToNCallback\n");
// The operand of `test.consumer` converts 1:N (i32 -> (i16, i16)), but the
// pattern only provides the 1:1 matchAndRewrite (matchAndRewrite1ToN is
// null). The driver's 1:1 dispatch cannot represent a 1:N-remapped operand,
// so the pattern fails to match and the conversion fails.
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.consumer\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAdd1ToNConversion(converter, splitI32, NULL);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
// Only the 1:1 matchAndRewrite is set; matchAndRewrite1ToN is null.
MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertConsumer,
NULL};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.consumer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsFailure(result) &&
"1:1-only pattern must fail to legalize a 1:N-remapped operand");
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverter1ToNOperandRequires1ToNCallback: PASSED
fprintf(stderr, "testTypeConverter1ToNOperandRequires1ToNCallback: PASSED\n");
}
void testTypeConverter1ToNConversionErasure(MlirContext ctx) {
// CHECK-LABEL: @testTypeConverter1ToNConversionErasure
fprintf(stderr, "@testTypeConverter1ToNConversionErasure\n");
// The i32 operand of `test.consumer` is erased (converted to zero types), so
// the 1:N pattern is invoked with no operands and builds a nullary
// `test.consumer_legal`.
const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
"\"test.consumer\"(%0) : (i32) -> ()\n";
MlirModule module =
mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
MlirOperation moduleOp = mlirModuleGetOperation(module);
MlirTypeConverter converter = mlirTypeConverterCreate();
mlirTypeConverterAdd1ToNConversion(converter, eraseI32, NULL);
MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
MlirConversionPatternCallbacks callbacks = {NULL, NULL, NULL,
convertConsumer1ToN};
MlirConversionPattern pattern = mlirOpConversionPatternCreate(
mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
callbacks, NULL, 0, NULL);
mlirRewritePatternSetAdd(patterns,
mlirConversionPatternAsRewritePattern(pattern));
MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
mlirRewritePatternSetDestroy(patterns);
MlirConversionTarget target = mlirConversionTargetCreate(ctx);
mlirConversionTargetAddIllegalOp(
target, mlirStringRefCreateFromCString("test.consumer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.producer"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("test.consumer_legal"));
mlirConversionTargetAddLegalOp(
target, mlirStringRefCreateFromCString("builtin.module"));
MlirConversionConfig config = mlirConversionConfigCreate();
MlirLogicalResult result =
mlirApplyPartialConversion(moduleOp, target, frozen, config);
assert(mlirLogicalResultIsSuccess(result));
mlirOperationDump(moduleOp);
// clang-format off
// CHECK: "test.consumer_legal"() : () -> ()
// clang-format on
mlirConversionConfigDestroy(config);
mlirConversionTargetDestroy(target);
mlirFrozenRewritePatternSetDestroy(frozen);
mlirTypeConverterDestroy(converter);
mlirModuleDestroy(module);
// CHECK: testTypeConverter1ToNConversionErasure: PASSED
fprintf(stderr, "testTypeConverter1ToNConversionErasure: PASSED\n");
}
int main(void) {
MlirContext ctx = mlirContextCreate();
mlirContextSetAllowUnregisteredDialects(ctx, true);
mlirContextGetOrLoadDialect(ctx, mlirStringRefCreateFromCString("builtin"));
testInsertionPoint(ctx);
testCreateBlock(ctx);
testInlineRegionBlock(ctx);
testReplaceOp(ctx);
testErase(ctx);
testMove(ctx);
testOpModification(ctx);
testReplaceUses(ctx);
testGreedyRewriteDriverConfig(ctx);
testCloneWithMapping(ctx);
testInsertionPointSaveRestore(ctx);
testConversionTargetDynamicLegality(ctx);
testTypeConverterSourceMaterialization(ctx);
testTypeConverterTargetMaterialization(ctx);
testTypeConverterConversionStatus(ctx);
testTypeConverter1ToNTargetMaterialization(ctx);
testTypeConverter1ToNConversionDeclineRollback(ctx);
testTypeConverter1ToNConversionFailure(ctx);
testTypeConverter1ToNTargetMaterializationDecline(ctx);
testTypeConverterMultiInputSourceMaterialization(ctx);
testConversionReplaceOpWithMultipleRanges(ctx);
testTypeConverter1ToNOperandRequires1ToNCallback(ctx);
testTypeConverter1ToNConversionErasure(ctx);
mlirContextDestroy(ctx);
return 0;
}