LLZK 3.0.0
An open-source IR for Zero Knowledge (ZK) circuits
Loading...
Searching...
No Matches
EmptyTemplateRemovalPass.cpp
Go to the documentation of this file.
1//===-- EmptyTemplateRemovalPass.cpp ----------------------------*- C++ -*-===//
2//
3// Part of the LLZK Project, under the Apache License v2.0.
4// See LICENSE.txt for license information.
5// Copyright 2025 Veridise Inc.
6// SPDX-License-Identifier: Apache-2.0
7//
8//===----------------------------------------------------------------------===//
13//===----------------------------------------------------------------------===//
14
23
24#include <mlir/Dialect/SCF/Transforms/Patterns.h>
25#include <mlir/Transforms/DialectConversion.h>
26
27// Include the generated base pass class definitions.
29#define GEN_PASS_DEF_EMPTYTEMPLATEREMOVALPASS
31} // namespace llzk::polymorphic
33#include "SharedImpl.h"
35#define DEBUG_TYPE "llzk-drop-empty-templates"
37using namespace mlir;
38using namespace llzk::array;
39using namespace llzk::component;
40using namespace llzk::function;
41using namespace llzk::pod;
42using namespace llzk::polymorphic;
44
45namespace {
46
47static inline bool hasEmptyParamList(StructType t) {
48 if (ArrayAttr paramList = t.getParams()) {
49 return paramList.empty();
50 }
51 return false;
52}
53
54/// Convert StructType with empty parameter list to one with no parameters.
55class EmptyParamListStructTypeConverter : public TypeConverter {
56public:
57 EmptyParamListStructTypeConverter() : TypeConverter() {
58
59 addConversion([](Type inputTy) { return inputTy; });
60
61 addConversion([](StructType inputTy) -> StructType {
62 return hasEmptyParamList(inputTy) ? StructType::get(inputTy.getNameRef()) : inputTy;
63 });
65 addConversion([this](ArrayType inputTy) {
66 // Recursively convert element type
67 return ArrayType::get(
68 this->convertType(inputTy.getElementType()), inputTy.getDimensionSizes()
69 );
70 });
71
72 addConversion([this](PodType inputTy) {
73 // Recursively convert record types
74 llvm::ArrayRef<RecordAttr> records = inputTy.getRecords();
75 if (records.empty()) {
76 return inputTy;
77 }
78 llvm::SmallVector<RecordAttr> newRecords;
79 newRecords.reserve(records.size());
80 MLIRContext *ctx = inputTy.getContext();
81 for (RecordAttr attr : records) {
82 newRecords.push_back(
83 RecordAttr::get(ctx, attr.getName(), this->convertType(attr.getType()))
84 );
85 }
86 return PodType::get(ctx, newRecords);
87 });
88 }
89};
90
91
92class DeleteNoDefTemplatePattern : public OpConversionPattern<TemplateOp> {
93public:
94 using OpConversionPattern<TemplateOp>::OpConversionPattern;
95
96 static inline bool legal(TemplateOp op) {
97 return llvm::any_of(op.getBodyRegion().getOps(), [](Operation &p) {
98 return llvm::isa<StructDefOp, FuncDefOp>(p);
99 });
100 }
101
102 LogicalResult matchAndRewrite(
103 TemplateOp op, TemplateOpAdaptor, ConversionPatternRewriter &rewriter
104 ) const override {
105 if (legal(op)) {
106 return failure();
107 }
108 LLVM_DEBUG({
109 llvm::dbgs() << "found template with no struct or function definitions: " << op << '\n';
110 });
111 rewriter.eraseOp(op);
112 return success();
113 }
114};
115
117class ReplaceNoParamTemplatePattern : public OpConversionPattern<TemplateOp> {
118public:
119 using OpConversionPattern<TemplateOp>::OpConversionPattern;
120
121 static inline bool legal(TemplateOp op) {
122 return op.hasConstOps<TemplateSymbolBindingOpInterface>();
123 }
124
125 LogicalResult matchAndRewrite(
126 TemplateOp op, TemplateOpAdaptor adaptor, ConversionPatternRewriter &rewriter
127 ) const override {
128 if (legal(op)) {
129 return failure();
130 }
131 LLVM_DEBUG({
132 llvm::dbgs() << "found template with no constant parameters or expressions: " << op << '\n';
133 });
134 // Convert types within the current body.
135 Region &currentBody = adaptor.getBodyRegion();
136 if (failed(rewriter.convertRegionTypes(&currentBody, *getTypeConverter()))) {
137 LLVM_DEBUG(llvm::dbgs() << "convertRegionTypes(currentBody) failed!\n");
138 return failure();
139 }
140 // Insert new ModuleOp at location of the current template.
141 ModuleOp newOp = rewriter.create<ModuleOp>(op.getLoc(), adaptor.getSymName());
142 // Move the current body into the module and erase the now-empty template op.
143 // First, clear body region of the new module to prepare for `inlineRegionBefore`.
144 Region &newOpBody = newOp.getBodyRegion();
145 if (!newOpBody.empty()) {
146 rewriter.eraseBlock(&newOpBody.front());
147 }
148 rewriter.inlineRegionBefore(currentBody, newOpBody, newOpBody.end());
149 rewriter.eraseOp(op);
150 return success();
151 }
152};
153
154class PassImpl : public llzk::polymorphic::impl::EmptyTemplateRemovalPassBase<PassImpl> {
155 using Base = EmptyTemplateRemovalPassBase<PassImpl>;
156 using Base::Base;
157
158 void runOnOperation() override {
159 ModuleOp modOp = getOperation();
160 MLIRContext *ctx = modOp.getContext();
161 EmptyParamListStructTypeConverter tyConv;
162 ConversionTarget target = newConverterDefinedTarget<>(tyConv, ctx);
163 // Mark TemplateOp legal only if legal according to both patterns.
164 target.addDynamicallyLegalOp<TemplateOp>([](TemplateOp op) {
165 return DeleteNoDefTemplatePattern::legal(op) && ReplaceNoParamTemplatePattern::legal(op);
166 });
167 RewritePatternSet patterns = llzk::newGeneralRewritePatternSet(tyConv, ctx, target);
168 // Try `DeleteNoDefTemplatePattern` first since full removal is better that replacement.
169 patterns.add<DeleteNoDefTemplatePattern>(tyConv, ctx);
170 patterns.add<ReplaceNoParamTemplatePattern>(tyConv, ctx);
171 if (failed(applyFullConversion(modOp, target, std::move(patterns)))) {
172 signalPassFailure();
173 }
174 }
175};
176
177} // namespace
Common private implementation for poly dialect passes.
::mlir::Type getElementType() const
static ArrayType get(::mlir::Type elementType, ::llvm::ArrayRef<::mlir::Attribute > dimensionSizes)
Definition Types.cpp.inc:83
::llvm::ArrayRef<::mlir::Attribute > getDimensionSizes() const
::mlir::SymbolRefAttr getNameRef() const
static StructType get(::mlir::SymbolRefAttr structName)
Definition Types.cpp.inc:79
::mlir::ArrayAttr getParams() const
static PodType get(::mlir::MLIRContext *context, ::llvm::ArrayRef<::llzk::pod::RecordAttr > records)
Definition Types.cpp.inc:68
::llvm::ArrayRef<::llzk::pod::RecordAttr > getRecords() const
::mlir::Region & getBodyRegion()
Definition Ops.h.inc:873
bool hasConstOps()
Return true if there are ops of type OpT within the body region.
Definition Ops.h.inc:929
mlir::ConversionTarget newConverterDefinedTarget(mlir::TypeConverter &tyConv, mlir::MLIRContext *ctx, AdditionalChecks &&...checks)
Return a new ConversionTarget allowing all LLZK-required dialects and defining Op legality based on t...
Definition SharedImpl.h:183
mlir::RewritePatternSet newGeneralRewritePatternSet(mlir::TypeConverter &tyConv, mlir::MLIRContext *ctx, mlir::ConversionTarget &target)
Return a new RewritePatternSet covering all LLZK op types that may contain a StructType.