diff --git a/apps/onnx/onnx_converter.cc b/apps/onnx/onnx_converter.cc index 75559585d088..8668142928f7 100644 --- a/apps/onnx/onnx_converter.cc +++ b/apps/onnx/onnx_converter.cc @@ -55,11 +55,11 @@ Halide::Expr inline_func_call(Halide::Expr e) { FuncCallInliner inliner; Halide::Expr r_old = Halide::Internal::simplify(e); - Halide::Expr r = inliner.mutate(r_old); + Halide::Expr r = inliner(r_old); while (!r.same_as(r_old)) { r_old = Halide::Internal::simplify(r); - r = inliner.mutate(r_old); + r = inliner(r_old); } return r; diff --git a/src/AddAtomicMutex.cpp b/src/AddAtomicMutex.cpp index 35ff29ba1330..b2243725d489 100644 --- a/src/AddAtomicMutex.cpp +++ b/src/AddAtomicMutex.cpp @@ -415,7 +415,7 @@ class AddAtomicMutex : public IRMutator { std::string name = unique_name('t'); index_let = index; index = Variable::make(index.type(), name); - body = ReplaceStoreIndexWithVar(op->producer_name, index).mutate(body); + body = ReplaceStoreIndexWithVar(op->producer_name, index)(body); } // This generates a pointer to the mutex array Expr mutex_array = Variable::make( @@ -454,8 +454,8 @@ Stmt add_atomic_mutex(Stmt s, const std::vector &outputs) { CheckAtomicValidity check; s.accept(&check); if (check.any_atomic) { - s = RemoveUnnecessaryMutexUse().mutate(s); - s = AddAtomicMutex(outputs).mutate(s); + s = RemoveUnnecessaryMutexUse()(s); + s = AddAtomicMutex(outputs)(s); } return s; } diff --git a/src/AddImageChecks.cpp b/src/AddImageChecks.cpp index 7114ad360135..b9ef94564049 100644 --- a/src/AddImageChecks.cpp +++ b/src/AddImageChecks.cpp @@ -103,6 +103,7 @@ class TrimStmtToPartsThatAccessBuffers : public IRMutator { bool touches_buffer = false; const map &buffers; +protected: using IRMutator::visit; Expr visit(const Call *op) override { @@ -185,10 +186,10 @@ Stmt add_image_checks_inner(Stmt s, // Add the input buffer(s) and annotate which output buffers are // used on host. - s.accept(&finder); + finder(s); Scope empty_scope; - Stmt sub_stmt = TrimStmtToPartsThatAccessBuffers(bufs).mutate(s); + Stmt sub_stmt = TrimStmtToPartsThatAccessBuffers(bufs)(s); map boxes = boxes_touched(sub_stmt, empty_scope, fb); // Now iterate through all the buffers, creating a list of lets @@ -225,7 +226,7 @@ Stmt add_image_checks_inner(Stmt s, string extent_name = concat_strings(name, ".extent.", i); string stride_name = concat_strings(name, ".stride.", i); replace_with_required[min_name] = Variable::make(Int(32), min_name + ".required"); - replace_with_required[extent_name] = simplify(Variable::make(Int(32), extent_name + ".required")); + replace_with_required[extent_name] = Variable::make(Int(32), extent_name + ".required"); replace_with_required[stride_name] = Variable::make(Int(32), stride_name + ".required"); } } @@ -737,6 +738,7 @@ Stmt add_image_checks(const Stmt &s, // Checks for images go at the marker deposited by computation // bounds inference. class Injector : public IRMutator { + protected: using IRMutator::visit; Expr visit(const Variable *op) override { @@ -794,9 +796,10 @@ Stmt add_image_checks(const Stmt &s, bool will_inject_host_copies) : outputs(outputs), t(t), order(order), env(env), fb(fb), will_inject_host_copies(will_inject_host_copies) { } - } injector(outputs, t, order, env, fb, will_inject_host_copies); + }; + Injector injector(outputs, t, order, env, fb, will_inject_host_copies); - return injector.mutate(s); + return injector(s); } } // namespace Internal diff --git a/src/AlignLoads.cpp b/src/AlignLoads.cpp index 263c6b4844de..9801411341dd 100644 --- a/src/AlignLoads.cpp +++ b/src/AlignLoads.cpp @@ -165,7 +165,7 @@ class AlignLoads : public IRMutator { } // namespace Stmt align_loads(const Stmt &s, int alignment, int min_bytes_to_align) { - return AlignLoads(alignment, min_bytes_to_align).mutate(s); + return AlignLoads(alignment, min_bytes_to_align)(s); } } // namespace Internal diff --git a/src/AllocationBoundsInference.cpp b/src/AllocationBoundsInference.cpp index a1e0831b975e..5598f286d276 100644 --- a/src/AllocationBoundsInference.cpp +++ b/src/AllocationBoundsInference.cpp @@ -169,8 +169,8 @@ class StripDeclareBoxTouched : public IRMutator { Stmt allocation_bounds_inference(Stmt s, const map &env, const FuncValueBounds &fb) { - s = AllocationInference(env, fb).mutate(s); - s = StripDeclareBoxTouched().mutate(s); + s = AllocationInference(env, fb)(s); + s = StripDeclareBoxTouched()(s); return s; } diff --git a/src/Associativity.cpp b/src/Associativity.cpp index 421fff6278f3..2351b79c88ab 100644 --- a/src/Associativity.cpp +++ b/src/Associativity.cpp @@ -339,7 +339,7 @@ AssociativeOp prove_associativity(const string &f, vector args, vector sema; std::set producers_dropped; @@ -285,6 +286,7 @@ class GenerateProducerBody : public NoOpCollapsingMutator { }; class GenerateConsumerBody : public NoOpCollapsingMutator { +protected: const string &func; vector sema; @@ -342,6 +344,7 @@ class GenerateConsumerBody : public NoOpCollapsingMutator { }; class CloneAcquire : public IRMutator { +protected: using IRMutator::visit; const string &old_name; @@ -390,6 +393,7 @@ class CountConsumeNodes : public IRVisitor { }; class ForkAsyncProducers : public IRMutator { +protected: using IRMutator::visit; const map &env; @@ -414,8 +418,8 @@ class ForkAsyncProducers : public IRMutator { sema_vars.push_back(Variable::make(type_of(), sema_names.back())); } - Stmt producer = GenerateProducerBody(name, sema_vars, cloned_acquires).mutate(body); - Stmt consumer = GenerateConsumerBody(name, sema_vars).mutate(body); + Stmt producer = GenerateProducerBody(name, sema_vars, cloned_acquires)(body); + Stmt consumer = GenerateConsumerBody(name, sema_vars)(body); // Recurse on both sides producer = mutate(producer); @@ -434,7 +438,7 @@ class ForkAsyncProducers : public IRMutator { // of the producer and consumer. const vector &clones = cloned_acquires[sema_name]; for (const auto &i : clones) { - body = CloneAcquire(sema_name, i).mutate(body); + body = CloneAcquire(sema_name, i)(body); body = LetStmt::make(i, sema_space, body); } @@ -493,6 +497,7 @@ class ForkAsyncProducers : public IRMutator { // simple failure case, error_async_require_fail. One has not been // written for the complex nested case yet.) class InitializeSemaphores : public IRMutator { +protected: using IRMutator::visit; const Type sema_type = type_of(); @@ -558,6 +563,7 @@ class InitializeSemaphores : public IRMutator { // A class to support stmt_uses_vars queries that repeatedly hit the same // sub-stmts. Used to support TightenProducerConsumerNodes below. class CachingStmtUsesVars : public IRMutator { +protected: const Scope<> &query; bool found_use = false; std::map cache; @@ -613,6 +619,7 @@ class CachingStmtUsesVars : public IRMutator { // Tighten the scope of consume nodes as much as possible to avoid needless synchronization. class TightenProducerConsumerNodes : public IRMutator { +protected: using IRMutator::visit; Stmt make_producer_consumer(const string &name, bool is_producer, Stmt body, const Scope<> &scope, CachingStmtUsesVars &uses_vars) { @@ -703,6 +710,7 @@ class TightenProducerConsumerNodes : public IRMutator { // Update indices to add ring buffer. class UpdateIndices : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const Provide *op) override { @@ -734,6 +742,7 @@ class UpdateIndices : public IRMutator { // Inject ring buffering. class InjectRingBuffering : public IRMutator { +protected: using IRMutator::visit; struct Loop { @@ -768,7 +777,7 @@ class InjectRingBuffering : public IRMutator { } current_index = current_index % f.schedule().ring_buffer(); // Adds an extra index for to the all of the references of f. - body = UpdateIndices(op->name, current_index).mutate(body); + body = UpdateIndices(op->name, current_index)(body); if (f.schedule().async()) { Expr sema_var = Variable::make(type_of(), f.name() + ".folding_semaphore.ring_buffer"); @@ -816,6 +825,7 @@ class InjectRingBuffering : public IRMutator { // Broaden the scope of acquire nodes to pack trailing work into the // same task and to potentially reduce the nesting depth of tasks. class ExpandAcquireNodes : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const Block *op) override { @@ -918,6 +928,7 @@ class ExpandAcquireNodes : public IRMutator { }; class TightenForkNodes : public IRMutator { +protected: using IRMutator::visit; Stmt make_fork(const Stmt &first, const Stmt &rest) { @@ -1005,12 +1016,12 @@ class TightenForkNodes : public IRMutator { } // namespace Stmt fork_async_producers(Stmt s, const map &env) { - s = TightenProducerConsumerNodes(env).mutate(s); - s = InjectRingBuffering(env).mutate(s); - s = ForkAsyncProducers(env).mutate(s); - s = ExpandAcquireNodes().mutate(s); - s = TightenForkNodes().mutate(s); - s = InitializeSemaphores().mutate(s); + s = TightenProducerConsumerNodes(env)(s); + s = InjectRingBuffering(env)(s); + s = ForkAsyncProducers(env)(s); + s = ExpandAcquireNodes()(s); + s = TightenForkNodes()(s); + s = InitializeSemaphores()(s); return s; } diff --git a/src/AutoScheduleUtils.cpp b/src/AutoScheduleUtils.cpp index d2227f831462..5f4578ee484f 100644 --- a/src/AutoScheduleUtils.cpp +++ b/src/AutoScheduleUtils.cpp @@ -53,14 +53,14 @@ Expr substitute_var_estimates(Expr e) { if (!e.defined()) { return e; } - return simplify(SubstituteVarEstimates().mutate(e)); + return simplify(SubstituteVarEstimates()(e)); } Stmt substitute_var_estimates(Stmt s) { if (!s.defined()) { return s; } - return simplify(SubstituteVarEstimates().mutate(s)); + return simplify(SubstituteVarEstimates()(s)); } int string_to_int(const string &s) { diff --git a/src/BoundConstantExtentLoops.cpp b/src/BoundConstantExtentLoops.cpp index c4c4a17eb297..bc76cbeb2738 100644 --- a/src/BoundConstantExtentLoops.cpp +++ b/src/BoundConstantExtentLoops.cpp @@ -12,6 +12,7 @@ namespace Internal { namespace { class BoundLoops : public IRMutator { +protected: using IRMutator::visit; std::vector> lets; @@ -128,7 +129,7 @@ class BoundLoops : public IRMutator { } // namespace Stmt bound_constant_extent_loops(const Stmt &s) { - return BoundLoops().mutate(s); + return BoundLoops()(s); } } // namespace Internal diff --git a/src/BoundSmallAllocations.cpp b/src/BoundSmallAllocations.cpp index f3347c0f47fd..18affb9236f5 100644 --- a/src/BoundSmallAllocations.cpp +++ b/src/BoundSmallAllocations.cpp @@ -156,7 +156,7 @@ class BoundSmallAllocations : public IRMutator { } // namespace Stmt bound_small_allocations(const Stmt &s) { - return BoundSmallAllocations().mutate(s); + return BoundSmallAllocations()(s); } } // namespace Internal diff --git a/src/Bounds.cpp b/src/Bounds.cpp index 0267485b0a79..f4493474c49f 100644 --- a/src/Bounds.cpp +++ b/src/Bounds.cpp @@ -231,7 +231,7 @@ class Bounds : public IRVisitor { #endif // DO_TRACK_BOUNDS_INTERVALS -private: +protected: // Compute the intrinsic bounds of a function. void bounds_of_func(const string &name, int value_index, Type t) { // if we can't get a good bound from the function, fall back to the bounds of the type. @@ -1799,7 +1799,7 @@ Interval bounds_of_expr_in_scope_with_indent(const Expr &expr, const Scope vars_depth; @@ -2255,7 +2256,7 @@ class BoxesTouched : public IRGraphVisitor { #endif // DO_TRACK_BOUNDS_INTERVALS -private: +protected: struct VarInstance { string var; int instance; @@ -3107,7 +3108,7 @@ map boxes_touched(const Expr &e, Stmt s, bool consider_calls, bool // as possible, so that BoxesTouched can prune the variable scope tighter // when encountering the IfThenElse. if (s.defined()) { - s = SolveIfThenElse().mutate(s); + s = SolveIfThenElse()(s); } // Do calls and provides separately, for better simplification. @@ -3116,18 +3117,18 @@ map boxes_touched(const Expr &e, Stmt s, bool consider_calls, bool if (consider_calls) { if (e.defined()) { - e.accept(&calls); + calls(e); } if (s.defined()) { - s.accept(&calls); + calls(s); } } if (consider_provides) { if (e.defined()) { - e.accept(&provides); + provides(e); } if (s.defined()) { - s.accept(&provides); + provides(s); } } diff --git a/src/BoundsInference.cpp b/src/BoundsInference.cpp index 72f45360b3b5..dda848e3bc7e 100644 --- a/src/BoundsInference.cpp +++ b/src/BoundsInference.cpp @@ -402,7 +402,7 @@ class BoundsInference : public IRMutator { } select_to_if_then_else; for (auto &e : exprs) { - e.value = select_to_if_then_else.mutate(e.value); + e.value = select_to_if_then_else(e.value); } } @@ -1382,8 +1382,7 @@ Stmt bounds_inference(Stmt s, s = For::make("", 0, 0, ForType::Serial, Partition::Never, DeviceAPI::None, s); s = BoundsInference(funcs, fused_func_groups, fused_pairs_in_groups, - outputs, func_bounds, target) - .mutate(s); + outputs, func_bounds, target)(s); return s.as()->body; } diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 036b92651667..1df6fcd6f548 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -519,6 +519,11 @@ target_compile_definitions(Halide PRIVATE WITH_SPIRV) target_compile_definitions(Halide PRIVATE WITH_VULKAN) target_compile_definitions(Halide PRIVATE WITH_WEBGPU) +if (WITH_COMPILER_PROFILING) + target_compile_definitions(Halide PRIVATE WITH_COMPILER_PROFILING) +endif() + + ## # Flatbuffers and Serialization dependencies. ## diff --git a/src/CSE.cpp b/src/CSE.cpp index c2a46d93bc4d..f4da1cd45ccc 100644 --- a/src/CSE.cpp +++ b/src/CSE.cpp @@ -186,6 +186,7 @@ class Replacer : public IRGraphMutator { }; class RemoveLets : public IRGraphMutator { +protected: using IRGraphMutator::visit; Scope scope; @@ -218,6 +219,7 @@ class RemoveLets : public IRGraphMutator { }; class CSEEveryExprInStmt : public IRMutator { +protected: bool lift_all; using IRMutator::visit; @@ -269,7 +271,7 @@ Expr common_subexpression_elimination(const Expr &e_in, bool lift_all) { debug(4) << "\n\n\nInput to CSE " << e << "\n"; - e = RemoveLets().mutate(e); + e = RemoveLets()(e); debug(4) << "After removing lets: " << e << "\n"; @@ -277,6 +279,7 @@ Expr common_subexpression_elimination(const Expr &e_in, bool lift_all) { // the same name as the temporaries we intend to introduce. Find any such // Vars so that we know not to use those names. class UniqueNameProvider : public IRGraphVisitor { + protected: using IRGraphVisitor::visit; const char prefix = 't'; // Annoyingly, this can't be static because this is a local class. @@ -303,14 +306,17 @@ Expr common_subexpression_elimination(const Expr &e_in, bool lift_all) { } while (vars.count(name)); return name; } - } namer; - e.accept(&namer); + }; + UniqueNameProvider namer; + { + e.accept(&namer); + } GVN gvn; - e = gvn.mutate(e); + e = gvn(e); ComputeUseCounts count_uses(gvn, lift_all); - count_uses.include(e); + count_uses(e); debug(4) << "Canonical form without lets " << e << "\n"; @@ -331,7 +337,7 @@ Expr common_subexpression_elimination(const Expr &e_in, bool lift_all) { // Rebuild the expr to include references to the variables: Replacer replacer(replacements); - e = replacer.mutate(e); + e = replacer(e); debug(4) << "With variables " << e << "\n"; @@ -340,7 +346,7 @@ Expr common_subexpression_elimination(const Expr &e_in, bool lift_all) { // Drop this variable as an acceptable replacement for this expr. replacer.erase(value); // Use containing lets in the value. - e = Let::make(var, replacer.mutate(value), e); + e = Let::make(var, replacer(value), e); } debug(4) << "With lets: " << e << "\n"; @@ -349,7 +355,7 @@ Expr common_subexpression_elimination(const Expr &e_in, bool lift_all) { } Stmt common_subexpression_elimination(const Stmt &s, bool lift_all) { - return CSEEveryExprInStmt(lift_all).mutate(s); + return CSEEveryExprInStmt(lift_all)(s); } // Testing code. @@ -388,8 +394,7 @@ class NormalizeVarNames : public IRMutator { void check(const Expr &in, const Expr &correct) { Expr result = common_subexpression_elimination(in); - NormalizeVarNames n; - result = n.mutate(result); + result = NormalizeVarNames()(result); internal_assert(equal(result, correct)) << "Incorrect CSE:\n" << in diff --git a/src/CanonicalizeGPUVars.cpp b/src/CanonicalizeGPUVars.cpp index 7ca9b7c4fbf5..609323f8e5dd 100644 --- a/src/CanonicalizeGPUVars.cpp +++ b/src/CanonicalizeGPUVars.cpp @@ -363,10 +363,8 @@ class ValidateGPUSchedule : public IRVisitor { } // anonymous namespace Stmt canonicalize_gpu_vars(Stmt s) { - ValidateGPUSchedule validator; - s.accept(&validator); - CanonicalizeGPUVars canonicalizer; - s = canonicalizer.mutate(s); + ValidateGPUSchedule()(s); + s = CanonicalizeGPUVars()(s); return s; } diff --git a/src/ClampUnsafeAccesses.cpp b/src/ClampUnsafeAccesses.cpp index ed6955446196..b976e55d8d0f 100644 --- a/src/ClampUnsafeAccesses.cpp +++ b/src/ClampUnsafeAccesses.cpp @@ -107,7 +107,7 @@ struct ClampUnsafeAccesses : IRMutator { } // namespace Stmt clamp_unsafe_accesses(const Stmt &s, const std::map &env, FuncValueBounds &func_bounds) { - return ClampUnsafeAccesses(env, func_bounds).mutate(s); + return ClampUnsafeAccesses(env, func_bounds)(s); } } // namespace Halide::Internal diff --git a/src/CodeGen_ARM.cpp b/src/CodeGen_ARM.cpp index 7178e82965d8..645763fe6a7f 100644 --- a/src/CodeGen_ARM.cpp +++ b/src/CodeGen_ARM.cpp @@ -1196,7 +1196,7 @@ void CodeGen_ARM::compile_func(const LoweredFunc &f, // Substitute in strided loads to get vld2/3/4 emission. We don't do it // on Apple silicon, because doing a dense load and then shuffling is // actually faster. - func.body = SubstituteInStridedLoads().mutate(func.body); + func.body = SubstituteInStridedLoads()(func.body); } // Look for opportunities to turn a + (b << c) into umlal/smlal // and a - (b << c) into umlsl/smlsl. diff --git a/src/CodeGen_D3D12Compute_Dev.cpp b/src/CodeGen_D3D12Compute_Dev.cpp index 7029193cc7c1..8385fb8b5448 100644 --- a/src/CodeGen_D3D12Compute_Dev.cpp +++ b/src/CodeGen_D3D12Compute_Dev.cpp @@ -1120,7 +1120,7 @@ void CodeGen_D3D12Compute_Dev::CodeGen_D3D12Compute_C::add_kernel(Stmt s, }; FindSharedAllocationsAndUniquify fsa; - s = fsa.mutate(s); + s = fsa(s); uint32_t total_shared_bytes = 0; for (const Stmt &sop : fsa.allocs) { diff --git a/src/CodeGen_Hexagon.cpp b/src/CodeGen_Hexagon.cpp index 78d1cbe0ed47..25f7e11aa885 100644 --- a/src/CodeGen_Hexagon.cpp +++ b/src/CodeGen_Hexagon.cpp @@ -316,7 +316,7 @@ class SloppyUnpredicateLoadsAndStores : public IRMutator { }; Stmt sloppy_unpredicate_loads_and_stores(const Stmt &s) { - return SloppyUnpredicateLoadsAndStores().mutate(s); + return SloppyUnpredicateLoadsAndStores()(s); } class InjectHVXLocks : public IRMutator { @@ -463,7 +463,7 @@ class InjectHVXLocks : public IRMutator { Stmt inject_hvx_lock_unlock(Stmt body, const Target &target) { InjectHVXLocks i(target); - body = i.mutate(body); + body = i(body); if (i.uses_hvx) { body = acquire_hvx_context(body, target); } diff --git a/src/CodeGen_Metal_Dev.cpp b/src/CodeGen_Metal_Dev.cpp index bac293b6ef16..61fbb8602c10 100644 --- a/src/CodeGen_Metal_Dev.cpp +++ b/src/CodeGen_Metal_Dev.cpp @@ -676,7 +676,6 @@ struct BufferSize { void CodeGen_Metal_Dev::CodeGen_Metal_C::add_kernel(const Stmt &s, const string &name, const vector &args) { - debug(2) << "Adding Metal kernel " << name << "\n"; // Figure out which arguments should be passed in constant. diff --git a/src/CodeGen_PTX_Dev.cpp b/src/CodeGen_PTX_Dev.cpp index 77783fd528aa..8f9fbb774f36 100644 --- a/src/CodeGen_PTX_Dev.cpp +++ b/src/CodeGen_PTX_Dev.cpp @@ -538,7 +538,7 @@ void CodeGen_PTX_Dev::codegen_vector_reduce(const VectorReduce *op, const Expr & Expr b_slice = Shuffle::make_slice(b, i + l * factor, 1, p.factor); i_slice = Call::make(i_slice.type(), p.name, {a_slice, b_slice, i_slice}, Call::PureExtern); } - i_slice = RewriteLoadsAs32Bit().mutate(i_slice); + i_slice = RewriteLoadsAs32Bit()(i_slice); i_slice = simplify(i_slice); i_slice = common_subexpression_elimination(i_slice); result.push_back(i_slice); diff --git a/src/CompilerLogger.cpp b/src/CompilerLogger.cpp index c22561a77a6e..58fe2f1760da 100644 --- a/src/CompilerLogger.cpp +++ b/src/CompilerLogger.cpp @@ -123,7 +123,7 @@ void JSONCompilerLogger::obfuscate() { std::string rule = it.first; for (const auto &e : it.second) { ObfuscateNames obfuscater; - n[rule].emplace_back(obfuscater.mutate(e)); + n[rule].emplace_back(obfuscater(e)); } } matched_simplifier_rules = n; @@ -136,8 +136,8 @@ void JSONCompilerLogger::obfuscate() { // to post-process output from multiple unrelated Generators // and combine Exprs with similar shapes. ObfuscateNames obfuscater; - auto failed_to_prove = obfuscater.mutate(it.first); - auto original_expr = obfuscater.mutate(it.second); + auto failed_to_prove = obfuscater(it.first); + auto original_expr = obfuscater(it.second); n.emplace_back(std::move(failed_to_prove), std::move(original_expr)); } failed_to_prove_exprs = n; diff --git a/src/DebugToFile.cpp b/src/DebugToFile.cpp index 89ea6c36c92d..d4f90a75b4fd 100644 --- a/src/DebugToFile.cpp +++ b/src/DebugToFile.cpp @@ -125,12 +125,12 @@ class AddDummyRealizations : public IRMutator { Stmt debug_to_file(Stmt s, const vector &outputs, const map &env) { // Temporarily wrap the produce nodes for the output functions in // realize nodes so that we know when to write the debug outputs. - s = AddDummyRealizations(outputs).mutate(s); + s = AddDummyRealizations(outputs)(s); - s = DebugToFile(env).mutate(s); + s = DebugToFile(env)(s); // Remove the realize node we wrapped around the output - s = RemoveDummyRealizations(outputs).mutate(s); + s = RemoveDummyRealizations(outputs)(s); return s; } diff --git a/src/Definition.cpp b/src/Definition.cpp index 5e6b00d95867..1cd83ce4d45b 100644 --- a/src/Definition.cpp +++ b/src/Definition.cpp @@ -6,6 +6,7 @@ #include "IR.h" #include "IRMutator.h" #include "IROperator.h" +#include "IRVisitor.h" #include "Var.h" namespace Halide { @@ -26,47 +27,47 @@ struct DefinitionContents { : predicate(const_true()) { } - void accept(IRVisitor *visitor) const { + void accept(IRVisitor &visitor) const { if (predicate.defined()) { - predicate.accept(visitor); + visitor(predicate); } for (const Expr &val : values) { - val.accept(visitor); + visitor(val); } for (const Expr &arg : args) { - arg.accept(visitor); + visitor(arg); } - stage_schedule.accept(visitor); + stage_schedule.accept(&visitor); for (const Specialization &s : specializations) { if (s.condition.defined()) { - s.condition.accept(visitor); + s.condition.accept(&visitor); } - s.definition.accept(visitor); + s.definition.accept(&visitor); } } - void mutate(IRMutator *mutator) { + void mutate(IRMutator &mutator) { if (predicate.defined()) { - predicate = mutator->mutate(predicate); + predicate = mutator(predicate); } for (auto &value : values) { - value = mutator->mutate(value); + value = mutator(value); } for (auto &arg : args) { - arg = mutator->mutate(arg); + arg = mutator(arg); } - stage_schedule.mutate(mutator); + stage_schedule.mutate(&mutator); for (Specialization &s : specializations) { if (s.condition.defined()) { - s.condition = mutator->mutate(s.condition); + s.condition = mutator(s.condition); } - s.definition.mutate(mutator); + s.definition.mutate(&mutator); } } }; @@ -146,11 +147,11 @@ bool Definition::is_init() const { } void Definition::accept(IRVisitor *visitor) const { - contents->accept(visitor); + contents->accept(*visitor); } void Definition::mutate(IRMutator *mutator) { - contents->mutate(mutator); + contents->mutate(*mutator); } std::vector &Definition::args() { diff --git a/src/Deinterleave.cpp b/src/Deinterleave.cpp index a1c881ea9744..6cf9428d9b14 100644 --- a/src/Deinterleave.cpp +++ b/src/Deinterleave.cpp @@ -172,8 +172,7 @@ class StoreCollector : public IRMutator { Stmt collect_strided_stores(const Stmt &stmt, const std::string &name, int stride, int max_stores, std::vector lets, std::vector &stores) { - StoreCollector collect(name, stride, max_stores, lets, stores); - return collect.mutate(stmt); + return StoreCollector(name, stride, max_stores, lets, stores)(stmt); } class Deinterleaver : public IRGraphMutator { @@ -407,8 +406,7 @@ class Deinterleaver : public IRGraphMutator { Expr deinterleave(Expr e, int starting_lane, int lane_stride, int new_lanes, const Scope<> &lets) { e = substitute_in_all_lets(e); - Deinterleaver d(starting_lane, lane_stride, new_lanes, lets); - e = d.mutate(e); + e = Deinterleaver(starting_lane, lane_stride, new_lanes, lets)(e); e = common_subexpression_elimination(e); return e; } @@ -802,7 +800,7 @@ class Interleaver : public IRMutator { } // namespace Stmt rewrite_interleavings(const Stmt &s) { - return Interleaver().mutate(s); + return Interleaver()(s); } namespace { diff --git a/src/DerivativeUtils.cpp b/src/DerivativeUtils.cpp index 23643010855f..54e1c37e0b84 100644 --- a/src/DerivativeUtils.cpp +++ b/src/DerivativeUtils.cpp @@ -116,7 +116,7 @@ Expr add_let_expression(const Expr &expr, const map &let_var_mapping, const vector &let_variables) { // TODO: find a faster way to do this - Expr ret = StripLets().mutate(expr); + Expr ret = StripLets()(expr); bool changed = true; vector injected(let_variables.size(), false); while (changed) { @@ -593,7 +593,7 @@ struct SubstituteCallArgWithPureArg : public IRMutator { } // namespace Expr substitute_call_arg_with_pure_arg(Func f, int variable_id, const Expr &e) { - return simplify(SubstituteCallArgWithPureArg(std::move(f), variable_id).mutate(e)); + return simplify(SubstituteCallArgWithPureArg(std::move(f), variable_id)(e)); } } // namespace Internal diff --git a/src/DistributeShifts.cpp b/src/DistributeShifts.cpp index 4d053e7d8dfe..1169466f947c 100644 --- a/src/DistributeShifts.cpp +++ b/src/DistributeShifts.cpp @@ -196,7 +196,7 @@ class DistributeShiftsAsMuls : public IRMutator { } // namespace Stmt distribute_shifts(const Stmt &s, bool multiply_adds) { - return DistributeShiftsAsMuls(multiply_adds).mutate(s); + return DistributeShiftsAsMuls(multiply_adds)(s); } } // namespace Internal diff --git a/src/EarlyFree.cpp b/src/EarlyFree.cpp index 8b664c2bcf8d..0014915a0700 100644 --- a/src/EarlyFree.cpp +++ b/src/EarlyFree.cpp @@ -1,11 +1,8 @@ -#include #include #include "EarlyFree.h" -#include "ExprUsesVar.h" -#include "IREquality.h" #include "IRMutator.h" -#include "InjectHostDevBufferCopies.h" +#include "IRVisitor.h" namespace Halide { namespace Internal { @@ -159,7 +156,7 @@ class InjectEarlyFrees : public IRMutator { InjectMarker inject_marker; inject_marker.func = alloc->name; inject_marker.last_use = last_use.last_use; - stmt = inject_marker.mutate(stmt); + stmt = inject_marker(stmt); } else { stmt = Allocate::make(alloc->name, alloc->type, alloc->memory_type, alloc->extents, alloc->condition, @@ -174,7 +171,7 @@ class InjectEarlyFrees : public IRMutator { Stmt inject_early_frees(const Stmt &s) { InjectEarlyFrees early_frees; - return early_frees.mutate(s); + return early_frees(s); } } // namespace Internal diff --git a/src/EliminateBoolVectors.cpp b/src/EliminateBoolVectors.cpp index e4afa8f21569..b68e1efa2d83 100644 --- a/src/EliminateBoolVectors.cpp +++ b/src/EliminateBoolVectors.cpp @@ -322,11 +322,11 @@ class EliminateBoolVectors : public IRMutator { } // namespace Stmt eliminate_bool_vectors(const Stmt &s) { - return EliminateBoolVectors().mutate(s); + return EliminateBoolVectors()(s); } Expr eliminate_bool_vectors(const Expr &e) { - return EliminateBoolVectors().mutate(e); + return EliminateBoolVectors()(e); } } // namespace Internal diff --git a/src/Expr.cpp b/src/Expr.cpp index d73bd72660fa..7d55fe9350c4 100644 --- a/src/Expr.cpp +++ b/src/Expr.cpp @@ -4,6 +4,20 @@ namespace Halide { namespace Internal { +const char *IRNodeType_string(IRNodeType type) { + switch (type) { +#define PROFILE_NODE_CASE(T) \ + case Halide::Internal::IRNodeType::T: \ + return #T; + + HALIDE_FOR_EACH_IR_NODE(PROFILE_NODE_CASE) +#undef PROFILE_NODE_CASE + + default: + internal_error << "Unknown Node Tag"; + } +} + const IntImm *IntImm::make(Type t, int64_t value) { internal_assert(t.is_int() && t.is_scalar()) << "IntImm must be a scalar Int\n"; diff --git a/src/Expr.h b/src/Expr.h index b9832c104de8..159f541c9b82 100644 --- a/src/Expr.h +++ b/src/Expr.h @@ -21,63 +21,76 @@ namespace Internal { class IRMutator; class IRVisitor; +// Exprs, in order of strength. Code in IRMatch.h and the +// simplifier relies on this order for canonicalization of +// expressions, so you may need to update those modules if you +// change this list. +#define HALIDE_FOR_EACH_IR_EXPR(X) \ + X(IntImm) \ + X(UIntImm) \ + X(FloatImm) \ + X(StringImm) \ + X(Broadcast) \ + X(Cast) \ + X(Reinterpret) \ + X(Variable) \ + X(Add) \ + X(Sub) \ + X(Mod) \ + X(Mul) \ + X(Div) \ + X(Min) \ + X(Max) \ + X(EQ) \ + X(NE) \ + X(LT) \ + X(LE) \ + X(GT) \ + X(GE) \ + X(And) \ + X(Or) \ + X(Not) \ + X(Select) \ + X(Load) \ + X(Ramp) \ + X(Call) \ + X(Let) \ + X(Shuffle) \ + X(VectorReduce) + +/* Stmts */ +#define HALIDE_FOR_EACH_IR_STMT(X) \ + X(LetStmt) \ + X(AssertStmt) \ + X(ProducerConsumer) \ + X(For) \ + X(Acquire) \ + X(Store) \ + X(Provide) \ + X(Allocate) \ + X(Free) \ + X(Realize) \ + X(Block) \ + X(Fork) \ + X(IfThenElse) \ + X(Evaluate) \ + X(Prefetch) \ + X(Atomic) \ + X(HoistedStorage) + +#define HALIDE_FOR_EACH_IR_NODE(X) \ + HALIDE_FOR_EACH_IR_EXPR(X) \ + HALIDE_FOR_EACH_IR_STMT(X) + /** All our IR node types get unique IDs for the purposes of RTTI */ -enum class IRNodeType { - // Exprs, in order of strength. Code in IRMatch.h and the - // simplifier relies on this order for canonicalization of - // expressions, so you may need to update those modules if you - // change this list. - IntImm, - UIntImm, - FloatImm, - StringImm, - Broadcast, - Cast, - Reinterpret, - Variable, - Add, - Sub, - Mod, - Mul, - Div, - Min, - Max, - EQ, - NE, - LT, - LE, - GT, - GE, - And, - Or, - Not, - Select, - Load, - Ramp, - Call, - Let, - Shuffle, - VectorReduce, - // Stmts - LetStmt, - AssertStmt, - ProducerConsumer, - For, - Acquire, - Store, - Provide, - Allocate, - Free, - Realize, - Block, - Fork, - IfThenElse, - Evaluate, - Prefetch, - Atomic, - HoistedStorage +enum class IRNodeType : uint8_t { +#define DECL_ENUM(X) X, + HALIDE_FOR_EACH_IR_NODE(DECL_ENUM) +#undef DECL_ENUM }; +const char *IRNodeType_string(IRNodeType type); + constexpr IRNodeType StrongestExprNodeType = IRNodeType::VectorReduce; /** The abstract base classes for a node in the Halide IR. */ diff --git a/src/ExtractTileOperations.cpp b/src/ExtractTileOperations.cpp index c9b234dc252f..cd315db389b4 100644 --- a/src/ExtractTileOperations.cpp +++ b/src/ExtractTileOperations.cpp @@ -672,7 +672,7 @@ class ExtractTileOperations : public IRMutator { } // namespace Stmt extract_tile_operations(const Stmt &s) { - return ExtractTileOperations().mutate(s); + return ExtractTileOperations()(s); } } // namespace Internal } // namespace Halide diff --git a/src/FindIntrinsics.cpp b/src/FindIntrinsics.cpp index f40d28644298..3aa495e5a621 100644 --- a/src/FindIntrinsics.cpp +++ b/src/FindIntrinsics.cpp @@ -1,17 +1,14 @@ #include "FindIntrinsics.h" #include "CSE.h" -#include "CodeGen_Internal.h" -#include "ConciseCasts.h" #include "ConstantBounds.h" #include "IRMatch.h" #include "IRMutator.h" +#include "IRVisitor.h" #include "Simplify.h" namespace Halide { namespace Internal { -using namespace Halide::ConciseCasts; - namespace { // This routine provides a guard on the return type of intrinsics such that only @@ -1121,6 +1118,7 @@ class FindIntrinsics : public IRMutator { // because each let in a chain has a wider value than the // ones it refers to. class SubstituteInWideningLets : public IRMutator { +protected: using IRMutator::visit; bool widens(const Expr &e) { @@ -1220,7 +1218,7 @@ class SubstituteInWideningLets : public IRMutator { if (should_replace) { size_t start_of_new_lets = frames.size(); - value = extractor.mutate(value); + value = extractor(value); // Mutate any subexpressions the extractor decided to // leave behind, in case they in turn depend on lets // we've decided to substitute in. @@ -1266,16 +1264,16 @@ class SubstituteInWideningLets : public IRMutator { } // namespace Stmt find_intrinsics(const Stmt &s) { - Stmt stmt = SubstituteInWideningLets().mutate(s); - stmt = FindIntrinsics().mutate(stmt); + Stmt stmt = SubstituteInWideningLets()(s); + stmt = FindIntrinsics()(stmt); // In case we want to hoist widening ops back out stmt = common_subexpression_elimination(stmt); return stmt; } Expr find_intrinsics(const Expr &e) { - Expr expr = SubstituteInWideningLets().mutate(e); - expr = FindIntrinsics().mutate(expr); + Expr expr = SubstituteInWideningLets()(e); + expr = FindIntrinsics()(expr); expr = common_subexpression_elimination(expr); return expr; } @@ -1651,11 +1649,11 @@ class LowerIntrinsics : public IRMutator { } // namespace Expr lower_intrinsics(const Expr &e) { - return LowerIntrinsics().mutate(e); + return LowerIntrinsics()(e); } Stmt lower_intrinsics(const Stmt &s) { - return LowerIntrinsics().mutate(s); + return LowerIntrinsics()(s); } } // namespace Internal diff --git a/src/FlattenNestedRamps.cpp b/src/FlattenNestedRamps.cpp index a98ed9bdb427..efa373f6970a 100644 --- a/src/FlattenNestedRamps.cpp +++ b/src/FlattenNestedRamps.cpp @@ -148,11 +148,11 @@ class LowerConcatBits : public IRMutator { } // namespace Stmt flatten_nested_ramps(const Stmt &s) { - return LowerConcatBits().mutate(FlattenRamps().mutate(s)); + return LowerConcatBits()(FlattenRamps()(s)); } Expr flatten_nested_ramps(const Expr &e) { - return LowerConcatBits().mutate(FlattenRamps().mutate(e)); + return LowerConcatBits()(FlattenRamps()(e)); } } // namespace Internal diff --git a/src/Func.cpp b/src/Func.cpp index 04b50412b6ac..a081116762df 100644 --- a/src/Func.cpp +++ b/src/Func.cpp @@ -614,7 +614,7 @@ vector substitute_self_reference(const vector &values, const string vector result; result.reserve(values.size()); for (const auto &val : values) { - result.push_back(subs.mutate(val)); + result.push_back(subs(val)); } return result; } diff --git a/src/Function.cpp b/src/Function.cpp index 54fe96f785f2..d9484e5aca0d 100644 --- a/src/Function.cpp +++ b/src/Function.cpp @@ -1,7 +1,5 @@ #include #include -#include -#include #include #include "CSE.h" @@ -168,10 +166,10 @@ struct FunctionContents { if (!extern_function_name.empty()) { for (ExternFuncArgument &i : extern_arguments) { if (i.is_expr()) { - i.expr = mutator->mutate(i.expr); + i.expr = (*mutator)(i.expr); } } - extern_proxy_expr = mutator->mutate(extern_proxy_expr); + extern_proxy_expr = (*mutator)(extern_proxy_expr); } } }; @@ -837,14 +835,14 @@ void Function::define_update(const vector &_args, vector values, con // memory leaks. We need to break these cycles. WeakenFunctionPtrs weakener(contents.get()); for (auto &arg : args) { - arg = weakener.mutate(arg); + arg = weakener(arg); } for (auto &value : values) { - value = weakener.mutate(value); + value = weakener(value); } if (check.reduction_domain.defined()) { check.reduction_domain.set_predicate( - weakener.mutate(check.reduction_domain.predicate())); + weakener(check.reduction_domain.predicate())); } Definition r(args, values, check.reduction_domain, false); diff --git a/src/FuseGPUThreadLoops.cpp b/src/FuseGPUThreadLoops.cpp index 88f9a542550f..ec85e5a383f9 100644 --- a/src/FuseGPUThreadLoops.cpp +++ b/src/FuseGPUThreadLoops.cpp @@ -31,6 +31,7 @@ using std::vector; namespace { class ExtractBlockSize : public IRVisitor { +protected: Expr block_extent[3], block_count[3]; string block_var_name[3]; @@ -123,6 +124,7 @@ class ExtractBlockSize : public IRVisitor { }; class NormalizeDimensionality : public IRMutator { +protected: using IRMutator::visit; const ExtractBlockSize &block_size; @@ -184,6 +186,7 @@ class NormalizeDimensionality : public IRMutator { }; class ReplaceForWithIf : public IRMutator { +protected: using IRMutator::visit; const ExtractBlockSize &block_size; @@ -223,6 +226,7 @@ class ReplaceForWithIf : public IRMutator { }; class ExtractSharedAndHeapAllocations : public IRMutator { +protected: using IRMutator::visit; struct IntInterval { @@ -289,7 +293,7 @@ class ExtractSharedAndHeapAllocations : public IRMutator { public: vector allocations; -private: +protected: map shared; bool in_threads = false; @@ -907,7 +911,7 @@ class ExtractSharedAndHeapAllocations : public IRMutator { : alloc_name(alloc_name), cluster_name(cluster_name), offset(offset) { } } rewriter{alloc.name, name, offset}; - s = rewriter.mutate(s); + s = rewriter(s); } // Define the group offset in terms of the previous group in the cluster @@ -1041,6 +1045,7 @@ class ExtractSharedAndHeapAllocations : public IRMutator { // block. Should only be run after shared allocations have already // been extracted. class ExtractRegisterAllocations : public IRMutator { +protected: using IRMutator::visit; struct RegisterAllocation { @@ -1206,6 +1211,7 @@ class ExtractRegisterAllocations : public IRMutator { }; class InjectThreadBarriers : public IRMutator { +protected: bool in_threads = false, injected_barrier; using IRMutator::visit; @@ -1360,6 +1366,7 @@ class InjectThreadBarriers : public IRMutator { }; class FuseGPUThreadLoopsSingleKernel : public IRMutator { +protected: using IRMutator::visit; const ExtractBlockSize &block_size; ExtractSharedAndHeapAllocations &block_allocations; @@ -1373,7 +1380,7 @@ class FuseGPUThreadLoopsSingleKernel : public IRMutator { << body << "\n\n"; NormalizeDimensionality n(block_size, op->device_api); - body = n.mutate(body); + body = n(body); debug(3) << "Normalized dimensionality:\n" << body << "\n\n"; @@ -1382,7 +1389,7 @@ class FuseGPUThreadLoopsSingleKernel : public IRMutator { ExtractRegisterAllocations register_allocs; ForType innermost_loop_type = ForType::GPUThread; if (block_size.threads_dimensions()) { - body = register_allocs.mutate(body); + body = register_allocs(body); if (register_allocs.has_lane_loop) { innermost_loop_type = ForType::GPULane; } @@ -1394,14 +1401,14 @@ class FuseGPUThreadLoopsSingleKernel : public IRMutator { if (register_allocs.has_thread_loop) { // If there's no loop over threads, everything is already synchronous. InjectThreadBarriers i{block_allocations, register_allocs}; - body = i.mutate(body); + body = i(body); } debug(3) << "Injected synchronization:\n" << body << "\n\n"; ReplaceForWithIf f(block_size); - body = f.mutate(body); + body = f(body); debug(3) << "Replaced for with if:\n" << body << "\n\n"; @@ -1448,6 +1455,7 @@ class FuseGPUThreadLoopsSingleKernel : public IRMutator { }; class FuseGPUThreadLoops : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const For *op) override { @@ -1463,17 +1471,17 @@ class FuseGPUThreadLoops : public IRMutator { // Do the analysis of thread block size and shared memory // usage. ExtractBlockSize block_size; - Stmt loop = Stmt(op); - loop.accept(&block_size); + block_size(op); + Stmt loop(op); ExtractSharedAndHeapAllocations block_allocations(op->device_api); - loop = block_allocations.mutate(loop); + loop = block_allocations(loop); debug(3) << "Pulled out shared allocations:\n" << loop << "\n\n"; // Mutate the inside of the kernel - loop = FuseGPUThreadLoopsSingleKernel(block_size, block_allocations).mutate(loop); + loop = FuseGPUThreadLoopsSingleKernel(block_size, block_allocations)(loop); loop = block_allocations.rewrap_kernel_launch(loop, block_size, op->device_api); @@ -1485,6 +1493,7 @@ class FuseGPUThreadLoops : public IRMutator { }; class ZeroGPULoopMins : public IRMutator { +protected: bool in_non_glsl_gpu = false; using IRMutator::visit; @@ -1516,13 +1525,14 @@ class ZeroGPULoopMins : public IRMutator { // Also used by InjectImageIntrinsics Stmt zero_gpu_loop_mins(const Stmt &s) { - return ZeroGPULoopMins().mutate(s); + return ZeroGPULoopMins()(s); } namespace { // Find the inner most GPU block of a statement. class FindInnermostGPUBlock : public IRVisitor { +protected: using IRVisitor::visit; void visit(const For *op) override { @@ -1540,6 +1550,7 @@ class FindInnermostGPUBlock : public IRVisitor { // Given a condition and a loop, add the condition // to the loop body. class AddConditionToALoop : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const For *op) override { @@ -1562,6 +1573,7 @@ class AddConditionToALoop : public IRMutator { // Push if statements between GPU blocks through all GPU blocks. // Throw error if the if statement has an else clause. class NormalizeIfStatements : public IRMutator { +protected: using IRMutator::visit; bool inside_gpu_blocks = false; @@ -1579,10 +1591,10 @@ class NormalizeIfStatements : public IRMutator { return IRMutator::visit(op); } FindInnermostGPUBlock find; - op->accept(&find); + find(op); if (find.found_gpu_block != nullptr) { internal_assert(!op->else_case.defined()) << "Found an if statement with else case between two GPU blocks.\n"; - return AddConditionToALoop(op->condition, find.found_gpu_block).mutate(op->then_case); + return AddConditionToALoop(op->condition, find.found_gpu_block)(op->then_case); } return IRMutator::visit(op); } @@ -1594,9 +1606,9 @@ Stmt fuse_gpu_thread_loops(Stmt s) { // NormalizeIfStatements pushes the predicates between GPU blocks // into the innermost GPU block. FuseGPUThreadLoops would then // merge the predicate into the merged GPU thread. - s = NormalizeIfStatements().mutate(s); - s = FuseGPUThreadLoops().mutate(s); - s = ZeroGPULoopMins().mutate(s); + s = NormalizeIfStatements()(s); + s = FuseGPUThreadLoops()(s); + s = ZeroGPULoopMins()(s); return s; } diff --git a/src/FuzzFloatStores.cpp b/src/FuzzFloatStores.cpp index 35f311224cd1..511141889680 100644 --- a/src/FuzzFloatStores.cpp +++ b/src/FuzzFloatStores.cpp @@ -27,7 +27,7 @@ class FuzzFloatStores : public IRMutator { } // namespace Stmt fuzz_float_stores(const Stmt &s) { - return FuzzFloatStores().mutate(s); + return FuzzFloatStores()(s); } } // namespace Internal diff --git a/src/HexagonOffload.cpp b/src/HexagonOffload.cpp index ddab37ecc9b3..5cafec409579 100644 --- a/src/HexagonOffload.cpp +++ b/src/HexagonOffload.cpp @@ -693,7 +693,7 @@ class ReplaceParams : public IRMutator { }; Stmt replace_params(const Stmt &s, const std::map &replacements) { - return ReplaceParams(replacements).mutate(s); + return ReplaceParams(replacements)(s); } class InjectHexagonRpc : public IRMutator { diff --git a/src/HexagonOptimize.cpp b/src/HexagonOptimize.cpp index 125e965f1b57..745dd9a6e808 100644 --- a/src/HexagonOptimize.cpp +++ b/src/HexagonOptimize.cpp @@ -331,7 +331,7 @@ Expr apply_patterns(Expr x, const vector &patterns, const Target &targe } // Mutate the operands with the given mutator. for (Expr &op : matches) { - op = op_mutator->mutate(op); + op = (*op_mutator)(op); } x = replace_pattern(x, matches, p); @@ -2291,8 +2291,8 @@ Stmt optimize_hexagon_shuffles(const Stmt &s, int lut_alignment) { Stmt scatter_gather_generator(Stmt s) { // Generate vscatter-vgather instruction if target >= v65 s = substitute_in_all_lets(s); - s = ScatterGatherGenerator().mutate(s); - s = SyncronizationBarriers().mutate(s); + s = ScatterGatherGenerator()(s); + s = SyncronizationBarriers()(s); s = common_subexpression_elimination(s); return s; } @@ -2316,24 +2316,24 @@ Stmt optimize_hexagon_instructions(Stmt s, const Target &t) { // Pattern match VectorReduce IR node. Handle vector reduce instructions // before OptimizePatterns to prevent being mutated by patterns like // (v0 + v1 * c) -> add_mpy - s = VectorReducePatterns().mutate(s); + s = VectorReducePatterns()(s); debug(4) << "Hexagon: Lowering after VectorReducePatterns\n" << s << "\n"; // Peephole optimize for Hexagon instructions. These can generate // interleaves and deinterleaves alongside the HVX intrinsics. - s = OptimizePatterns(t).mutate(s); + s = OptimizePatterns(t)(s); debug(4) << "Hexagon: Lowering after OptimizePatterns\n" << s << "\n"; // Try to eliminate any redundant interleave/deinterleave pairs. - s = EliminateInterleaves(t, t.natural_vector_size(Int(8))).mutate(s); + s = EliminateInterleaves(t, t.natural_vector_size(Int(8)))(s); debug(4) << "Hexagon: Lowering after EliminateInterleaves\n" << s << "\n"; // There may be interleaves left over that we can fuse with other // operations. - s = FuseInterleaves().mutate(s); + s = FuseInterleaves()(s); debug(4) << "Hexagon: Lowering after FuseInterleaves\n" << s << "\n"; return s; diff --git a/src/IRMatch.cpp b/src/IRMatch.cpp index ffbb9406ad1e..7649fa9c2152 100644 --- a/src/IRMatch.cpp +++ b/src/IRMatch.cpp @@ -406,7 +406,7 @@ class WithLanes : public IRMutator { } // namespace Expr with_lanes(const Expr &x, int lanes) { - return WithLanes(lanes).mutate(x); + return WithLanes(lanes)(x); } } // namespace Internal diff --git a/src/IRMutator.cpp b/src/IRMutator.cpp index 9eecd0579840..c5ae276bda62 100644 --- a/src/IRMutator.cpp +++ b/src/IRMutator.cpp @@ -45,63 +45,59 @@ Expr IRMutator::visit(const Reinterpret *op) { return Reinterpret::make(op->type, std::move(value)); } -namespace { -template -Expr mutate_binary_operator(IRMutator *mutator, const T *op) { - Expr a = mutator->mutate(op->a); - Expr b = mutator->mutate(op->b); - if (a.same_as(op->a) && - b.same_as(op->b)) { - return op; - } - return T::make(std::move(a), std::move(b)); -} -} // namespace +#define mutate_binary_operator \ + Expr a = mutate(op->a); \ + Expr b = mutate(op->b); \ + if (a.same_as(op->a) && \ + b.same_as(op->b)) { \ + return op; \ + } \ + return std::decay_t::make(std::move(a), std::move(b)) Expr IRMutator::visit(const Add *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Sub *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Mul *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Div *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Mod *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Min *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Max *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const EQ *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const NE *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const LT *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const LE *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const GT *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const GE *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const And *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Or *op) { - return mutate_binary_operator(this, op); + mutate_binary_operator; } Expr IRMutator::visit(const Not *op) { diff --git a/src/IRMutator.h b/src/IRMutator.h index 148caf56a54b..5a389fc4870d 100644 --- a/src/IRMutator.h +++ b/src/IRMutator.h @@ -34,73 +34,48 @@ class IRMutator { * these in your subclass to mutate sub-expressions and * sub-statements. */ - virtual Expr mutate(const Expr &expr); - virtual Stmt mutate(const Stmt &stmt); + inline Expr operator()(const Expr &expr) { + return mutate(expr); + } + + inline Stmt operator()(const Stmt &stmt) { + return mutate(stmt); + } + + // Like mutate_with_changes, but discard the changes flag. + std::vector operator()(const std::vector &exprs) { + return mutate_with_changes(exprs).first; + } // Mutate all the Exprs and return the new list in ret, along with // a flag that is true iff at least one item in the list changed. std::pair, bool> mutate_with_changes(const std::vector &); +protected: + virtual Expr mutate(const Expr &expr); + virtual Stmt mutate(const Stmt &stmt); // Like mutate_with_changes, but discard the changes flag. std::vector mutate(const std::vector &exprs) { return mutate_with_changes(exprs).first; } -protected: // ExprNode<> and StmtNode<> are allowed to call visit (to implement mutate_expr/mutate_stmt()) template friend struct ExprNode; template friend struct StmtNode; - - virtual Expr visit(const IntImm *); - virtual Expr visit(const UIntImm *); - virtual Expr visit(const FloatImm *); - virtual Expr visit(const StringImm *); - virtual Expr visit(const Cast *); - virtual Expr visit(const Reinterpret *); - virtual Expr visit(const Add *); - virtual Expr visit(const Sub *); - virtual Expr visit(const Mul *); - virtual Expr visit(const Div *); - virtual Expr visit(const Mod *); - virtual Expr visit(const Min *); - virtual Expr visit(const Max *); - virtual Expr visit(const EQ *); - virtual Expr visit(const NE *); - virtual Expr visit(const LT *); - virtual Expr visit(const LE *); - virtual Expr visit(const GT *); - virtual Expr visit(const GE *); - virtual Expr visit(const And *); - virtual Expr visit(const Or *); - virtual Expr visit(const Not *); - virtual Expr visit(const Select *); - virtual Expr visit(const Load *); - virtual Expr visit(const Ramp *); - virtual Expr visit(const Broadcast *); - virtual Expr visit(const Let *); - virtual Stmt visit(const LetStmt *); - virtual Stmt visit(const AssertStmt *); - virtual Stmt visit(const ProducerConsumer *); - virtual Stmt visit(const Store *); - virtual Stmt visit(const Provide *); - virtual Stmt visit(const Allocate *); - virtual Stmt visit(const Free *); - virtual Stmt visit(const Realize *); - virtual Stmt visit(const Block *); - virtual Stmt visit(const Fork *); - virtual Stmt visit(const IfThenElse *); - virtual Stmt visit(const Evaluate *); - virtual Expr visit(const Call *); - virtual Expr visit(const Variable *); - virtual Stmt visit(const For *); - virtual Stmt visit(const Acquire *); - virtual Expr visit(const Shuffle *); - virtual Stmt visit(const Prefetch *); - virtual Stmt visit(const HoistedStorage *); - virtual Stmt visit(const Atomic *); - virtual Expr visit(const VectorReduce *); + template + friend std::pair mutate_region(Mutator *mutator, const Region &bounds, Args &&...args); + +#define HALIDE_DECL_VISIT_EXPR(T) \ + virtual Expr visit(const T *op); + HALIDE_FOR_EACH_IR_EXPR(HALIDE_DECL_VISIT_EXPR) +#undef HALIDE_DECL_VISIT_EXPR + +#define HALIDE_DECL_VISIT_STMT(T) \ + virtual Stmt visit(const T *op); + HALIDE_FOR_EACH_IR_STMT(HALIDE_DECL_VISIT_STMT) +#undef HALIDE_DECL_VISIT_STMT }; /** A mutator that caches and reapplies previously done mutations so @@ -111,10 +86,17 @@ class IRGraphMutator : public IRMutator { std::map expr_replacements; std::map stmt_replacements; -public: using IRMutator::mutate; Stmt mutate(const Stmt &s) override; Expr mutate(const Expr &e) override; + +public: + inline Expr operator()(const Expr &expr) { + return mutate(expr); + } + inline Stmt operator()(const Stmt &stmt) { + return mutate(stmt); + } }; /** A lambda-based IR mutator that accepts multiple lambdas for different @@ -131,6 +113,9 @@ struct LambdaMutator final : IRMutator { return IRMutator::visit(op); } +public: + using IRMutator::mutate; + private: LambdaOverloads handlers; @@ -149,150 +134,18 @@ struct LambdaMutator final : IRMutator { } protected: - Expr visit(const IntImm *op) override { - return this->visit_impl(op); - } - Expr visit(const UIntImm *op) override { - return this->visit_impl(op); - } - Expr visit(const FloatImm *op) override { - return this->visit_impl(op); - } - Expr visit(const StringImm *op) override { - return this->visit_impl(op); - } - Expr visit(const Cast *op) override { - return this->visit_impl(op); - } - Expr visit(const Reinterpret *op) override { - return this->visit_impl(op); - } - Expr visit(const Add *op) override { - return this->visit_impl(op); - } - Expr visit(const Sub *op) override { - return this->visit_impl(op); - } - Expr visit(const Mul *op) override { - return this->visit_impl(op); - } - Expr visit(const Div *op) override { - return this->visit_impl(op); - } - Expr visit(const Mod *op) override { - return this->visit_impl(op); - } - Expr visit(const Min *op) override { - return this->visit_impl(op); - } - Expr visit(const Max *op) override { - return this->visit_impl(op); - } - Expr visit(const EQ *op) override { - return this->visit_impl(op); - } - Expr visit(const NE *op) override { - return this->visit_impl(op); - } - Expr visit(const LT *op) override { - return this->visit_impl(op); - } - Expr visit(const LE *op) override { - return this->visit_impl(op); - } - Expr visit(const GT *op) override { - return this->visit_impl(op); - } - Expr visit(const GE *op) override { - return this->visit_impl(op); - } - Expr visit(const And *op) override { - return this->visit_impl(op); - } - Expr visit(const Or *op) override { - return this->visit_impl(op); - } - Expr visit(const Not *op) override { - return this->visit_impl(op); - } - Expr visit(const Select *op) override { - return this->visit_impl(op); - } - Expr visit(const Load *op) override { - return this->visit_impl(op); - } - Expr visit(const Ramp *op) override { - return this->visit_impl(op); - } - Expr visit(const Broadcast *op) override { - return this->visit_impl(op); - } - Expr visit(const Let *op) override { - return this->visit_impl(op); - } - Stmt visit(const LetStmt *op) override { - return this->visit_impl(op); - } - Stmt visit(const AssertStmt *op) override { - return this->visit_impl(op); - } - Stmt visit(const ProducerConsumer *op) override { - return this->visit_impl(op); - } - Stmt visit(const Store *op) override { - return this->visit_impl(op); - } - Stmt visit(const Provide *op) override { - return this->visit_impl(op); - } - Stmt visit(const Allocate *op) override { - return this->visit_impl(op); - } - Stmt visit(const Free *op) override { - return this->visit_impl(op); - } - Stmt visit(const Realize *op) override { - return this->visit_impl(op); - } - Stmt visit(const Block *op) override { - return this->visit_impl(op); - } - Stmt visit(const Fork *op) override { - return this->visit_impl(op); - } - Stmt visit(const IfThenElse *op) override { - return this->visit_impl(op); - } - Stmt visit(const Evaluate *op) override { - return this->visit_impl(op); - } - Expr visit(const Call *op) override { - return this->visit_impl(op); - } - Expr visit(const Variable *op) override { - return this->visit_impl(op); - } - Stmt visit(const For *op) override { - return this->visit_impl(op); - } - Stmt visit(const Acquire *op) override { - return this->visit_impl(op); - } - Expr visit(const Shuffle *op) override { - return this->visit_impl(op); - } - Stmt visit(const Prefetch *op) override { - return this->visit_impl(op); - } - Stmt visit(const HoistedStorage *op) override { - return this->visit_impl(op); - } - Stmt visit(const Atomic *op) override { - return this->visit_impl(op); - } - Expr visit(const VectorReduce *op) override { - return this->visit_impl(op); - } +#define HALIDE_CALL_VISIT_EXPR_IMPL(T) \ + Expr visit(const T *op) override { \ + return this->visit_impl(op); \ + } + HALIDE_FOR_EACH_IR_EXPR(HALIDE_CALL_VISIT_EXPR_IMPL) +#undef HALIDE_CALL_VISIT_EXPR_IMPL +#define HALIDE_CALL_VISIT_STMT_IMPL(T) \ + Stmt visit(const T *op) override { \ + return this->visit_impl(op); \ + } + HALIDE_FOR_EACH_IR_STMT(HALIDE_CALL_VISIT_STMT_IMPL) +#undef HALIDE_CALL_VISIT_STMT_IMPL }; /** A lambda-based IR mutator that accepts multiple lambdas for overloading @@ -312,6 +165,7 @@ struct LambdaMutatorGeneric final : IRMutator { return IRMutator::mutate(op); } +public: Expr mutate(const Expr &e) override { if constexpr (std::is_invocable_v) { return handlers(this, e); @@ -338,7 +192,7 @@ auto mutate_with(const T &ir, Lambdas &&...lambdas) { using Generic = LambdaMutatorGeneric; if constexpr (std::is_invocable_v || std::is_invocable_v) { - return LambdaMutatorGeneric{std::forward(lambdas)...}.mutate(ir); + return LambdaMutatorGeneric{std::forward(lambdas)...}(ir); } else { LambdaMutator mutator{std::forward(lambdas)...}; // Each lambda must take two args: (auto *self, op). @@ -350,7 +204,7 @@ auto mutate_with(const T &ir, Lambdas &&...lambdas) { ...); static_assert(all_take_two_args, "All mutate_with lambdas must take two arguments: (auto *self, const T *op)"); - return mutator.mutate(ir); + return mutator(ir); } } diff --git a/src/IRPrinter.h b/src/IRPrinter.h index de812bdb0c47..12ffacb4fa88 100644 --- a/src/IRPrinter.h +++ b/src/IRPrinter.h @@ -65,7 +65,7 @@ class Closure; struct Interval; struct ConstantInterval; struct ModulusRemainder; -enum class IRNodeType; +enum class IRNodeType : uint8_t; /** Emit a halide node type on an output stream (such as std::cout) in * human-readable form */ diff --git a/src/IRVisitor.h b/src/IRVisitor.h index a14e71558e89..322f2bff6ba3 100644 --- a/src/IRVisitor.h +++ b/src/IRVisitor.h @@ -21,6 +21,19 @@ class IRVisitor { IRVisitor() = default; virtual ~IRVisitor() = default; + inline void operator()(const Stmt &s) { + s.accept(this); + } + + inline void operator()(const Expr &e) { + e.accept(this); + } + + template + inline void operator()(const T *op) { + visit(op); + } + protected: // ExprNode<> and StmtNode<> are allowed to call visit (to implement accept()) template @@ -29,54 +42,10 @@ class IRVisitor { template friend struct StmtNode; - virtual void visit(const IntImm *); - virtual void visit(const UIntImm *); - virtual void visit(const FloatImm *); - virtual void visit(const StringImm *); - virtual void visit(const Cast *); - virtual void visit(const Reinterpret *); - virtual void visit(const Add *); - virtual void visit(const Sub *); - virtual void visit(const Mul *); - virtual void visit(const Div *); - virtual void visit(const Mod *); - virtual void visit(const Min *); - virtual void visit(const Max *); - virtual void visit(const EQ *); - virtual void visit(const NE *); - virtual void visit(const LT *); - virtual void visit(const LE *); - virtual void visit(const GT *); - virtual void visit(const GE *); - virtual void visit(const And *); - virtual void visit(const Or *); - virtual void visit(const Not *); - virtual void visit(const Select *); - virtual void visit(const Load *); - virtual void visit(const Ramp *); - virtual void visit(const Broadcast *); - virtual void visit(const Let *); - virtual void visit(const LetStmt *); - virtual void visit(const AssertStmt *); - virtual void visit(const ProducerConsumer *); - virtual void visit(const Store *); - virtual void visit(const Provide *); - virtual void visit(const Allocate *); - virtual void visit(const Free *); - virtual void visit(const Realize *); - virtual void visit(const Block *); - virtual void visit(const Fork *); - virtual void visit(const IfThenElse *); - virtual void visit(const Evaluate *); - virtual void visit(const Call *); - virtual void visit(const Variable *); - virtual void visit(const For *); - virtual void visit(const Acquire *); - virtual void visit(const Shuffle *); - virtual void visit(const Prefetch *); - virtual void visit(const HoistedStorage *); - virtual void visit(const Atomic *); - virtual void visit(const VectorReduce *); +#define HALIDE_DECL_VISIT(T) \ + virtual void visit(const T *op); + HALIDE_FOR_EACH_IR_NODE(HALIDE_DECL_VISIT) +#undef HALIDE_DECL_VISIT }; /** A lambda-based IR visitor that accepts multiple lambdas for different @@ -111,150 +80,18 @@ struct LambdaVisitor final : IRVisitor { } protected: - void visit(const IntImm *op) override { - this->visit_impl(op); - } - void visit(const UIntImm *op) override { - this->visit_impl(op); - } - void visit(const FloatImm *op) override { - this->visit_impl(op); - } - void visit(const StringImm *op) override { - this->visit_impl(op); - } - void visit(const Cast *op) override { - this->visit_impl(op); - } - void visit(const Reinterpret *op) override { - this->visit_impl(op); - } - void visit(const Add *op) override { - this->visit_impl(op); - } - void visit(const Sub *op) override { - this->visit_impl(op); - } - void visit(const Mul *op) override { - this->visit_impl(op); - } - void visit(const Div *op) override { - this->visit_impl(op); - } - void visit(const Mod *op) override { - this->visit_impl(op); - } - void visit(const Min *op) override { - this->visit_impl(op); - } - void visit(const Max *op) override { - this->visit_impl(op); - } - void visit(const EQ *op) override { - this->visit_impl(op); - } - void visit(const NE *op) override { - this->visit_impl(op); - } - void visit(const LT *op) override { - this->visit_impl(op); - } - void visit(const LE *op) override { - this->visit_impl(op); - } - void visit(const GT *op) override { - this->visit_impl(op); - } - void visit(const GE *op) override { - this->visit_impl(op); - } - void visit(const And *op) override { - this->visit_impl(op); - } - void visit(const Or *op) override { - this->visit_impl(op); - } - void visit(const Not *op) override { - this->visit_impl(op); - } - void visit(const Select *op) override { - this->visit_impl(op); - } - void visit(const Load *op) override { - this->visit_impl(op); - } - void visit(const Ramp *op) override { - this->visit_impl(op); - } - void visit(const Broadcast *op) override { - this->visit_impl(op); - } - void visit(const Let *op) override { - this->visit_impl(op); - } - void visit(const LetStmt *op) override { - this->visit_impl(op); - } - void visit(const AssertStmt *op) override { - this->visit_impl(op); - } - void visit(const ProducerConsumer *op) override { - this->visit_impl(op); - } - void visit(const Store *op) override { - this->visit_impl(op); - } - void visit(const Provide *op) override { - this->visit_impl(op); - } - void visit(const Allocate *op) override { - this->visit_impl(op); - } - void visit(const Free *op) override { - this->visit_impl(op); - } - void visit(const Realize *op) override { - this->visit_impl(op); - } - void visit(const Block *op) override { - this->visit_impl(op); - } - void visit(const Fork *op) override { - this->visit_impl(op); - } - void visit(const IfThenElse *op) override { - this->visit_impl(op); - } - void visit(const Evaluate *op) override { - this->visit_impl(op); - } - void visit(const Call *op) override { - this->visit_impl(op); - } - void visit(const Variable *op) override { - this->visit_impl(op); - } - void visit(const For *op) override { - this->visit_impl(op); - } - void visit(const Acquire *op) override { - this->visit_impl(op); - } - void visit(const Shuffle *op) override { - this->visit_impl(op); - } - void visit(const Prefetch *op) override { - this->visit_impl(op); - } - void visit(const HoistedStorage *op) override { - this->visit_impl(op); - } - void visit(const Atomic *op) override { - this->visit_impl(op); - } - void visit(const VectorReduce *op) override { - this->visit_impl(op); - } +#define HALIDE_CALL_VISIT_EXPR_IMPL(T) \ + void visit(const T *op) override { \ + this->visit_impl(op); \ + } + HALIDE_FOR_EACH_IR_EXPR(HALIDE_CALL_VISIT_EXPR_IMPL) +#undef HALIDE_CALL_VISIT_EXPR_IMPL +#define HALIDE_CALL_VISIT_STMT_IMPL(T) \ + void visit(const T *op) override { \ + this->visit_impl(op); \ + } + HALIDE_FOR_EACH_IR_STMT(HALIDE_CALL_VISIT_STMT_IMPL) +#undef HALIDE_CALL_VISIT_STMT_IMPL }; template @@ -281,6 +118,20 @@ void visit_with(const IRHandle &ir, Lambdas &&...lambdas) { * without visiting the same node twice. This is for passes that are * capable of interpreting the IR as a DAG instead of a tree. */ class IRGraphVisitor : public IRVisitor { +public: + inline void operator()(const Expr &e) { + include(e); + } + inline void operator()(const Stmt &s) { + include(s); + } + +private: + /** The nodes visited so far. Only includes nodes with a ref count greater + * than one, because we know that nodes with a ref count of 1 will only be + * visited once if their parents are only visited once. */ + std::set visited; + protected: /** By default these methods add the node to the visited set, and * return whether or not it was already there. If it wasn't there, @@ -291,64 +142,13 @@ class IRGraphVisitor : public IRVisitor { virtual void include(const Stmt &); // @} -private: - /** The nodes visited so far. Only includes nodes with a ref count greater - * than one, because we know that nodes with a ref count of 1 will only be - * visited once if their parents are only visited once. */ - std::set visited; - -protected: /** These methods should call 'include' on the children to only * visit them if they haven't been visited already. */ // @{ - void visit(const IntImm *) override; - void visit(const UIntImm *) override; - void visit(const FloatImm *) override; - void visit(const StringImm *) override; - void visit(const Cast *) override; - void visit(const Reinterpret *) override; - void visit(const Add *) override; - void visit(const Sub *) override; - void visit(const Mul *) override; - void visit(const Div *) override; - void visit(const Mod *) override; - void visit(const Min *) override; - void visit(const Max *) override; - void visit(const EQ *) override; - void visit(const NE *) override; - void visit(const LT *) override; - void visit(const LE *) override; - void visit(const GT *) override; - void visit(const GE *) override; - void visit(const And *) override; - void visit(const Or *) override; - void visit(const Not *) override; - void visit(const Select *) override; - void visit(const Load *) override; - void visit(const Ramp *) override; - void visit(const Broadcast *) override; - void visit(const Let *) override; - void visit(const LetStmt *) override; - void visit(const AssertStmt *) override; - void visit(const ProducerConsumer *) override; - void visit(const Store *) override; - void visit(const Provide *) override; - void visit(const Allocate *) override; - void visit(const Free *) override; - void visit(const Realize *) override; - void visit(const Block *) override; - void visit(const Fork *) override; - void visit(const IfThenElse *) override; - void visit(const Evaluate *) override; - void visit(const Call *) override; - void visit(const Variable *) override; - void visit(const For *) override; - void visit(const Acquire *) override; - void visit(const Shuffle *) override; - void visit(const Prefetch *) override; - void visit(const HoistedStorage *) override; - void visit(const Atomic *) override; - void visit(const VectorReduce *) override; +#define HALIDE_VISIT_OVERRIDE(T) \ + void visit(const T *) override; + HALIDE_FOR_EACH_IR_NODE(HALIDE_VISIT_OVERRIDE) +#undef HALIDE_VISIT_OVERRIDE // @} }; @@ -366,88 +166,13 @@ class VariadicVisitor { return ExprRet{}; } switch (node->node_type) { - case IRNodeType::IntImm: - return ((T *)this)->visit((const IntImm *)node, std::forward(args)...); - case IRNodeType::UIntImm: - return ((T *)this)->visit((const UIntImm *)node, std::forward(args)...); - case IRNodeType::FloatImm: - return ((T *)this)->visit((const FloatImm *)node, std::forward(args)...); - case IRNodeType::StringImm: - return ((T *)this)->visit((const StringImm *)node, std::forward(args)...); - case IRNodeType::Broadcast: - return ((T *)this)->visit((const Broadcast *)node, std::forward(args)...); - case IRNodeType::Cast: - return ((T *)this)->visit((const Cast *)node, std::forward(args)...); - case IRNodeType::Reinterpret: - return ((T *)this)->visit((const Reinterpret *)node, std::forward(args)...); - case IRNodeType::Variable: - return ((T *)this)->visit((const Variable *)node, std::forward(args)...); - case IRNodeType::Add: - return ((T *)this)->visit((const Add *)node, std::forward(args)...); - case IRNodeType::Sub: - return ((T *)this)->visit((const Sub *)node, std::forward(args)...); - case IRNodeType::Mod: - return ((T *)this)->visit((const Mod *)node, std::forward(args)...); - case IRNodeType::Mul: - return ((T *)this)->visit((const Mul *)node, std::forward(args)...); - case IRNodeType::Div: - return ((T *)this)->visit((const Div *)node, std::forward(args)...); - case IRNodeType::Min: - return ((T *)this)->visit((const Min *)node, std::forward(args)...); - case IRNodeType::Max: - return ((T *)this)->visit((const Max *)node, std::forward(args)...); - case IRNodeType::EQ: - return ((T *)this)->visit((const EQ *)node, std::forward(args)...); - case IRNodeType::NE: - return ((T *)this)->visit((const NE *)node, std::forward(args)...); - case IRNodeType::LT: - return ((T *)this)->visit((const LT *)node, std::forward(args)...); - case IRNodeType::LE: - return ((T *)this)->visit((const LE *)node, std::forward(args)...); - case IRNodeType::GT: - return ((T *)this)->visit((const GT *)node, std::forward(args)...); - case IRNodeType::GE: - return ((T *)this)->visit((const GE *)node, std::forward(args)...); - case IRNodeType::And: - return ((T *)this)->visit((const And *)node, std::forward(args)...); - case IRNodeType::Or: - return ((T *)this)->visit((const Or *)node, std::forward(args)...); - case IRNodeType::Not: - return ((T *)this)->visit((const Not *)node, std::forward(args)...); - case IRNodeType::Select: - return ((T *)this)->visit((const Select *)node, std::forward(args)...); - case IRNodeType::Load: - return ((T *)this)->visit((const Load *)node, std::forward(args)...); - case IRNodeType::Ramp: - return ((T *)this)->visit((const Ramp *)node, std::forward(args)...); - case IRNodeType::Call: - return ((T *)this)->visit((const Call *)node, std::forward(args)...); - case IRNodeType::Let: - return ((T *)this)->visit((const Let *)node, std::forward(args)...); - case IRNodeType::Shuffle: - return ((T *)this)->visit((const Shuffle *)node, std::forward(args)...); - case IRNodeType::VectorReduce: - return ((T *)this)->visit((const VectorReduce *)node, std::forward(args)...); - // Explicitly list the Stmt types rather than using a - // default case so that when new IR nodes are added we - // don't miss them here. - case IRNodeType::LetStmt: - case IRNodeType::AssertStmt: - case IRNodeType::ProducerConsumer: - case IRNodeType::For: - case IRNodeType::Acquire: - case IRNodeType::Store: - case IRNodeType::Provide: - case IRNodeType::Allocate: - case IRNodeType::Free: - case IRNodeType::Realize: - case IRNodeType::Block: - case IRNodeType::Fork: - case IRNodeType::IfThenElse: - case IRNodeType::Evaluate: - case IRNodeType::Prefetch: - case IRNodeType::Atomic: - case IRNodeType::HoistedStorage: +#define HALIDE_SWITCH_EXPR(NT) \ + case IRNodeType::NT: \ + return ((T *)this)->visit((const NT *)node, std::forward(args)...); + HALIDE_FOR_EACH_IR_EXPR(HALIDE_SWITCH_EXPR) +#undef HALIDE_SWITCH_EXPR + + default: internal_error << "Unreachable"; } return ExprRet{}; @@ -459,73 +184,15 @@ class VariadicVisitor { return StmtRet{}; } switch (node->node_type) { - case IRNodeType::IntImm: - case IRNodeType::UIntImm: - case IRNodeType::FloatImm: - case IRNodeType::StringImm: - case IRNodeType::Broadcast: - case IRNodeType::Cast: - case IRNodeType::Reinterpret: - case IRNodeType::Variable: - case IRNodeType::Add: - case IRNodeType::Sub: - case IRNodeType::Mod: - case IRNodeType::Mul: - case IRNodeType::Div: - case IRNodeType::Min: - case IRNodeType::Max: - case IRNodeType::EQ: - case IRNodeType::NE: - case IRNodeType::LT: - case IRNodeType::LE: - case IRNodeType::GT: - case IRNodeType::GE: - case IRNodeType::And: - case IRNodeType::Or: - case IRNodeType::Not: - case IRNodeType::Select: - case IRNodeType::Load: - case IRNodeType::Ramp: - case IRNodeType::Call: - case IRNodeType::Let: - case IRNodeType::Shuffle: - case IRNodeType::VectorReduce: +#define HALIDE_SWITCH_STMT(NT) \ + case IRNodeType::NT: \ + return ((T *)this)->visit((const NT *)node, std::forward(args)...); + HALIDE_FOR_EACH_IR_STMT(HALIDE_SWITCH_STMT) +#undef HALIDE_SWITCH_STMT + + default: internal_error << "Unreachable"; break; - case IRNodeType::LetStmt: - return ((T *)this)->visit((const LetStmt *)node, std::forward(args)...); - case IRNodeType::AssertStmt: - return ((T *)this)->visit((const AssertStmt *)node, std::forward(args)...); - case IRNodeType::ProducerConsumer: - return ((T *)this)->visit((const ProducerConsumer *)node, std::forward(args)...); - case IRNodeType::For: - return ((T *)this)->visit((const For *)node, std::forward(args)...); - case IRNodeType::Acquire: - return ((T *)this)->visit((const Acquire *)node, std::forward(args)...); - case IRNodeType::Store: - return ((T *)this)->visit((const Store *)node, std::forward(args)...); - case IRNodeType::Provide: - return ((T *)this)->visit((const Provide *)node, std::forward(args)...); - case IRNodeType::Allocate: - return ((T *)this)->visit((const Allocate *)node, std::forward(args)...); - case IRNodeType::Free: - return ((T *)this)->visit((const Free *)node, std::forward(args)...); - case IRNodeType::Realize: - return ((T *)this)->visit((const Realize *)node, std::forward(args)...); - case IRNodeType::Block: - return ((T *)this)->visit((const Block *)node, std::forward(args)...); - case IRNodeType::Fork: - return ((T *)this)->visit((const Fork *)node, std::forward(args)...); - case IRNodeType::IfThenElse: - return ((T *)this)->visit((const IfThenElse *)node, std::forward(args)...); - case IRNodeType::Evaluate: - return ((T *)this)->visit((const Evaluate *)node, std::forward(args)...); - case IRNodeType::Prefetch: - return ((T *)this)->visit((const Prefetch *)node, std::forward(args)...); - case IRNodeType::Atomic: - return ((T *)this)->visit((const Atomic *)node, std::forward(args)...); - case IRNodeType::HoistedStorage: - return ((T *)this)->visit((const HoistedStorage *)node, std::forward(args)...); } return StmtRet{}; } diff --git a/src/InjectHostDevBufferCopies.cpp b/src/InjectHostDevBufferCopies.cpp index 4a899ee61fd6..6e8b21686657 100644 --- a/src/InjectHostDevBufferCopies.cpp +++ b/src/InjectHostDevBufferCopies.cpp @@ -29,6 +29,7 @@ Stmt call_extern_and_assert(const string &name, const vector &args) { namespace { class FindBufferUsage : public IRVisitor { +protected: using IRVisitor::visit; void visit(const Load *op) override { @@ -142,6 +143,7 @@ class FindBufferUsage : public IRVisitor { // the buffer as we go, sniffing usage within each leaf using // FindBufferUsage, and injecting device buffer logic as needed. class InjectBufferCopiesForSingleBuffer : public IRMutator { +protected: using IRMutator::visit; // The buffer being managed @@ -218,7 +220,7 @@ class InjectBufferCopiesForSingleBuffer : public IRMutator { // Sniff what happens to the buffer inside the stmt FindBufferUsage local_finder(buffer, DeviceAPI::Host); if (!precomputed) { - s.accept(&local_finder); + local_finder(s); precomputed = &local_finder; } FindBufferUsage &finder = *precomputed; @@ -356,7 +358,7 @@ class InjectBufferCopiesForSingleBuffer : public IRMutator { // leaf. Stmt visit(const For *op) override { FindBufferUsage finder(buffer, DeviceAPI::Host); - op->accept(&finder); + finder(op); if (finder.devices_touched.size() > 1) { // The state of the buffer going into the loop is the // union of the state before the loop starts and the state @@ -480,12 +482,14 @@ class InjectBufferCopiesForSingleBuffer : public IRMutator { // Inject the buffer-handling logic for all internal // allocations. Inputs and outputs are handled below. class InjectBufferCopies : public IRMutator { +protected: using IRMutator::visit; // Inject the registration of a device destructor just after the // .buffer symbol is defined (which is safely before the first // device_malloc). class InjectDeviceDestructor : public IRMutator { + protected: using IRMutator::visit; Stmt visit(const LetStmt *op) override { @@ -514,6 +518,7 @@ class InjectBufferCopies : public IRMutator { // and an Allocate node that takes its host field from the // .buffer. class InjectCombinedAllocation : public IRMutator { + protected: using IRMutator::visit; Stmt visit(const LetStmt *op) override { @@ -582,7 +587,7 @@ class InjectBufferCopies : public IRMutator { Stmt visit(const Allocate *op) override { FindBufferUsage finder(op->name, DeviceAPI::Host); - op->body.accept(&finder); + finder(op->body); bool touched_on_host = finder.devices_touched.count(DeviceAPI::Host); bool touched_on_device = finder.devices_touched.size() > (touched_on_host ? 1 : 0); @@ -595,7 +600,7 @@ class InjectBufferCopies : public IRMutator { Stmt body = mutate(op->body); InjectBufferCopiesForSingleBuffer injector(op->name, false, op->memory_type); - body = injector.mutate(body); + body = injector(body); string buffer_name = op->name + ".buffer"; Expr buffer = Variable::make(Handle(), buffer_name); @@ -621,8 +626,7 @@ class InjectBufferCopies : public IRMutator { Expr device_interface = make_device_interface_call(touching_device, op->memory_type); return InjectCombinedAllocation(op->name, op->type, op->extents, - op->condition, device_interface) - .mutate(body); + op->condition, device_interface)(body); } else { // Only touched on host but passed to an extern stage, or // only touched on device, or touched on multiple @@ -636,7 +640,7 @@ class InjectBufferCopies : public IRMutator { } // Add a device destructor - body = InjectDeviceDestructor(buffer_name).mutate(body); + body = InjectDeviceDestructor(buffer_name)(body); Expr condition = op->condition; bool touched_on_one_device = !touched_on_host && finder.devices_touched.size() == 1 && @@ -673,6 +677,7 @@ class InjectBufferCopies : public IRMutator { // ProducerConsumer node. Sometimes it's a Block containing a pair of // them. class FindOutermostProduce : public IRVisitor { +protected: using IRVisitor::visit; void visit(const Block *op) override { @@ -695,10 +700,12 @@ class FindOutermostProduce : public IRVisitor { // Inject the buffer handling code for the inputs and outputs at the // appropriate site. class InjectBufferCopiesForInputsAndOutputs : public IRMutator { +protected: Stmt site; // Find all references to external buffers. class FindInputsAndOutputs : public IRVisitor { + protected: using IRVisitor::visit; void include(const Parameter &p) { @@ -752,10 +759,10 @@ class InjectBufferCopiesForInputsAndOutputs : public IRMutator { Stmt mutate(const Stmt &s) override { if (s.same_as(site)) { FindInputsAndOutputs finder; - s.accept(&finder); + finder(s); Stmt new_stmt = s; for (const string &buf : finder.result) { - new_stmt = InjectBufferCopiesForSingleBuffer(buf, true, finder.result_storage.at(buf)).mutate(new_stmt); + new_stmt = InjectBufferCopiesForSingleBuffer(buf, true, finder.result_storage.at(buf))(new_stmt); } return new_stmt; } else { @@ -782,15 +789,15 @@ Stmt inject_host_dev_buffer_copies(Stmt s, const Target &t) { } // Handle internal allocations - s = InjectBufferCopies().mutate(s); + s = InjectBufferCopies()(s); // Handle inputs and outputs FindOutermostProduce outermost; - s.accept(&outermost); + outermost(s); if (outermost.result.defined()) { // If the entire pipeline simplified away, or just dispatches // to another pipeline, there may be no outermost produce. - s = InjectBufferCopiesForInputsAndOutputs(outermost.result).mutate(s); + s = InjectBufferCopiesForInputsAndOutputs(outermost.result)(s); } return s; diff --git a/src/Inline.cpp b/src/Inline.cpp index 31b6efcdf749..ce829f1c7326 100644 --- a/src/Inline.cpp +++ b/src/Inline.cpp @@ -181,15 +181,13 @@ class Inliner : public IRMutator { } }; -Stmt inline_function(Stmt s, const Function &f) { - Inliner i(f); - s = i.mutate(s); - return s; +Stmt inline_function(const Stmt &s, const Function &f) { + return Inliner(f)(s); } Expr inline_function(Expr e, const Function &f) { Inliner i(f); - e = i.mutate(e); + e = i(e); // TODO: making this > 1 should be desirable, // but explodes compiletimes in some situations. if (i.found > 0) { diff --git a/src/Inline.h b/src/Inline.h index fbbd78751e18..344e7c7ddf6d 100644 --- a/src/Inline.h +++ b/src/Inline.h @@ -16,7 +16,7 @@ class Function; * be inlined, it must not have any specializations (i.e. it can only have one * values definition). */ // @{ -Stmt inline_function(Stmt s, const Function &f); +Stmt inline_function(const Stmt &s, const Function &f); Expr inline_function(Expr e, const Function &f); void inline_function(Function caller, const Function &f); // @} diff --git a/src/InlineReductions.cpp b/src/InlineReductions.cpp index e3fc8c1311d2..2bd2873b2736 100644 --- a/src/InlineReductions.cpp +++ b/src/InlineReductions.cpp @@ -126,7 +126,7 @@ Expr sum(const RDom &r, Expr e, const Func &f) { << " passed to sum already has a definition"; Internal::FindFreeVars v(r, f.name()); - e = v.mutate(common_subexpression_elimination(e)); + e = v(common_subexpression_elimination(e)); user_assert(v.rdom.defined()) << "Expression passed to sum must reference a reduction domain"; @@ -152,7 +152,7 @@ Expr saturating_sum(const RDom &r, Expr e, const Func &f) { << " passed to saturating_sum already has a definition"; Internal::FindFreeVars v(r, f.name()); - e = v.mutate(common_subexpression_elimination(e)); + e = v(common_subexpression_elimination(e)); user_assert(v.rdom.defined()) << "Expression passed to saturating_sum must reference a reduction domain"; @@ -179,7 +179,7 @@ Expr product(const RDom &r, Expr e, const Func &f) { << " passed to product already has a definition"; Internal::FindFreeVars v(r, f.name()); - e = v.mutate(common_subexpression_elimination(e)); + e = v(common_subexpression_elimination(e)); user_assert(v.rdom.defined()) << "Expression passed to product must reference a reduction domain"; @@ -205,7 +205,7 @@ Expr maximum(const RDom &r, Expr e, const Func &f) { << " passed to maximum already has a definition"; Internal::FindFreeVars v(r, f.name()); - e = v.mutate(common_subexpression_elimination(e)); + e = v(common_subexpression_elimination(e)); user_assert(v.rdom.defined()) << "Expression passed to maximum must reference a reduction domain"; @@ -232,7 +232,7 @@ Expr minimum(const RDom &r, Expr e, const Func &f) { << " passed to minimum already has a definition"; Internal::FindFreeVars v(r, f.name()); - e = v.mutate(common_subexpression_elimination(e)); + e = v(common_subexpression_elimination(e)); user_assert(v.rdom.defined()) << "Expression passed to minimum must reference a reduction domain"; @@ -259,7 +259,7 @@ Tuple argmax(const RDom &r, Expr e, const Func &f) { << " passed to argmax already has a definition"; Internal::FindFreeVars v(r, f.name()); - e = v.mutate(common_subexpression_elimination(e)); + e = v(common_subexpression_elimination(e)); user_assert(v.rdom.defined()) << "Expression passed to argmax must reference a reduction domain"; @@ -298,7 +298,7 @@ Tuple argmin(const RDom &r, Expr e, const Func &f) { << " passed to argmin already has a definition"; Internal::FindFreeVars v(r, f.name()); - e = v.mutate(common_subexpression_elimination(e)); + e = v(common_subexpression_elimination(e)); user_assert(v.rdom.defined()) << "Expression passed to argmin must reference a reduction domain"; diff --git a/src/LICM.cpp b/src/LICM.cpp index ca5fcec3bd8b..d0ee80aad177 100644 --- a/src/LICM.cpp +++ b/src/LICM.cpp @@ -21,6 +21,7 @@ namespace { // Is it safe to lift an Expr out of a loop (and potentially across a device boundary) class CanLift : public IRVisitor { +protected: using IRVisitor::visit; void visit(const Call *op) override { @@ -54,6 +55,7 @@ class CanLift : public IRVisitor { // Lift pure loop invariants to the top level. Applied independently // to each loop. class LiftLoopInvariants : public IRMutator { +protected: using IRMutator::visit; Scope<> varying; @@ -183,6 +185,7 @@ class LiftLoopInvariants : public IRMutator { // them as just renamings of other variables. Easier to substitute // them in as a post-pass rather than make the pass above more clever. class SubstituteTrivialLets : public IRMutator { +protected: using IRMutator::visit; Expr visit(const Let *op) override { @@ -203,6 +206,7 @@ class SubstituteTrivialLets : public IRMutator { }; class LICM : public IRMutator { +protected: using IRMutator::visit; bool in_gpu_loop{false}; @@ -246,8 +250,8 @@ class LICM : public IRMutator { // Lift invariants LiftLoopInvariants lifter; - Stmt new_stmt = lifter.mutate(op); - new_stmt = SubstituteTrivialLets().mutate(new_stmt); + Stmt new_stmt = lifter(op); + new_stmt = SubstituteTrivialLets()(new_stmt); // As an optimization to reduce register pressure, take // the set of expressions to lift and check if any can @@ -336,6 +340,7 @@ class LICM : public IRMutator { // Reassociate summations to group together the loop invariants. Useful to run before LICM. class GroupLoopInvariants : public IRMutator { +protected: using IRMutator::visit; Scope var_depth; @@ -520,9 +525,9 @@ class GroupLoopInvariants : public IRMutator { } // namespace Stmt hoist_loop_invariant_values(Stmt s) { - s = GroupLoopInvariants().mutate(s); + s = GroupLoopInvariants()(s); s = common_subexpression_elimination(s); - s = LICM().mutate(s); + s = LICM()(s); s = simplify_exprs(s); return s; } @@ -532,6 +537,7 @@ namespace { // Move IfThenElse nodes from the inside of a piece of Stmt IR to the // outside when legal. class HoistIfStatements : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const LetStmt *op) override { @@ -656,9 +662,8 @@ class HoistIfStatements : public IRMutator { } // namespace -Stmt hoist_loop_invariant_if_statements(Stmt s) { - s = HoistIfStatements().mutate(s); - return s; +Stmt hoist_loop_invariant_if_statements(const Stmt &s) { + return HoistIfStatements()(s); } } // namespace Internal diff --git a/src/LICM.h b/src/LICM.h index 3d04db35143e..6c9cd16116d7 100644 --- a/src/LICM.h +++ b/src/LICM.h @@ -19,7 +19,7 @@ Stmt hoist_loop_invariant_values(Stmt); /** Just hoist loop-invariant if statements as far up as * possible. Does not lift other values. It's useful to run this * earlier in lowering to simplify the IR. */ -Stmt hoist_loop_invariant_if_statements(Stmt); +Stmt hoist_loop_invariant_if_statements(const Stmt &); } // namespace Internal } // namespace Halide diff --git a/src/LoopCarry.cpp b/src/LoopCarry.cpp index cc7fddb619da..51a4bbd03b7c 100644 --- a/src/LoopCarry.cpp +++ b/src/LoopCarry.cpp @@ -167,7 +167,7 @@ class StepForwards : public IRGraphMutator { Expr step_forwards(Expr e, const Scope &linear) { StepForwards step(linear); - e = step.mutate(e); + e = step(e); if (!step.success) { return Expr(); } else { @@ -554,7 +554,7 @@ class LoopCarry : public IRMutator { Stmt stmt; Stmt body = mutate(op->body); LoopCarryOverLoop carry(op->name, in_consume, max_carried_values); - body = carry.mutate(body); + body = carry(body); if (body.same_as(op->body)) { stmt = op; } else { @@ -583,9 +583,8 @@ class LoopCarry : public IRMutator { } // namespace -Stmt loop_carry(Stmt s, int max_carried_values) { - s = LoopCarry(max_carried_values).mutate(s); - return s; +Stmt loop_carry(const Stmt &s, int max_carried_values) { + return LoopCarry(max_carried_values)(s); } } // namespace Internal diff --git a/src/LoopCarry.h b/src/LoopCarry.h index f473e2627d54..e409935a0b3d 100644 --- a/src/LoopCarry.h +++ b/src/LoopCarry.h @@ -12,7 +12,7 @@ namespace Internal { * pessimization depending on how good the L1 cache is on the architecture * and how many memory issue slots there are. Currently only intended * for Hexagon. */ -Stmt loop_carry(Stmt, int max_carried_values = 8); +Stmt loop_carry(const Stmt &, int max_carried_values = 8); } // namespace Internal } // namespace Halide diff --git a/src/Lower.cpp b/src/Lower.cpp index 9b55bd20840d..da189b6f3a03 100644 --- a/src/Lower.cpp +++ b/src/Lower.cpp @@ -468,7 +468,7 @@ void lower_impl(const vector &output_funcs, if (!custom_passes.empty()) { for (size_t i = 0; i < custom_passes.size(); i++) { debug(1) << "Running custom lowering pass " << i << "...\n"; - s = custom_passes[i]->mutate(s); + s = (*custom_passes[i])(s); debug(1) << "Lowering after custom pass " << i << ":\n" << s << "\n\n"; } diff --git a/src/LowerParallelTasks.cpp b/src/LowerParallelTasks.cpp index 62e909136841..6d28d9779cbd 100644 --- a/src/LowerParallelTasks.cpp +++ b/src/LowerParallelTasks.cpp @@ -426,7 +426,7 @@ struct LowerParallelTasks : public IRMutator { Stmt lower_parallel_tasks(const Stmt &s, std::vector &closure_implementations, const std::string &name, const Target &t) { LowerParallelTasks lowering_mutator(name, t); - Stmt result = lowering_mutator.mutate(s); + Stmt result = lowering_mutator(s); // Main body will be dumped as part of standard lowering debugging, but closures will not be. debug(2) << [&] { diff --git a/src/LowerWarpShuffles.cpp b/src/LowerWarpShuffles.cpp index 7be557318bbb..9f244aad2ce0 100644 --- a/src/LowerWarpShuffles.cpp +++ b/src/LowerWarpShuffles.cpp @@ -785,8 +785,8 @@ class HoistWarpShuffles : public IRMutator { Stmt else_case = mutate(op->else_case); HoistWarpShufflesFromSingleIfStmt hoister; - then_case = hoister.mutate(then_case); - else_case = hoister.mutate(else_case); + then_case = hoister(then_case); + else_case = hoister(else_case); Stmt s = IfThenElse::make(op->condition, then_case, else_case); if (hoister.success) { return hoister.rewrap(s); @@ -794,7 +794,7 @@ class HoistWarpShuffles : public IRMutator { // Need to move the ifstmt further inwards instead. internal_assert(!else_case.defined()) << "Cannot hoist warp shuffle out of " << s << "\n"; string pred_name = unique_name('p'); - s = MoveIfStatementInwards(Variable::make(op->condition.type(), pred_name)).mutate(then_case); + s = MoveIfStatementInwards(Variable::make(op->condition.type(), pred_name))(then_case); return LetStmt::make(pred_name, op->condition, s); } } @@ -824,8 +824,8 @@ class LowerWarpShufflesInEachKernel : public IRMutator { Stmt visit(const For *op) override { if (op->device_api == DeviceAPI::CUDA && has_lane_loop(op)) { Stmt s = op; - s = LowerWarpShuffles(cuda_cap).mutate(s); - s = HoistWarpShuffles().mutate(s); + s = LowerWarpShuffles(cuda_cap)(s); + s = HoistWarpShuffles()(s); return simplify(s); } else { return IRMutator::visit(op); @@ -844,9 +844,9 @@ class LowerWarpShufflesInEachKernel : public IRMutator { Stmt lower_warp_shuffles(Stmt s, const Target &t) { s = hoist_loop_invariant_values(s); - s = SubstituteInLaneVar().mutate(s); + s = SubstituteInLaneVar()(s); s = simplify(s); - s = LowerWarpShufflesInEachKernel(t.get_cuda_capability_lower_bound()).mutate(s); + s = LowerWarpShufflesInEachKernel(t.get_cuda_capability_lower_bound())(s); return s; }; diff --git a/src/Memoization.cpp b/src/Memoization.cpp index 21cfbd4c9dce..b55c4a66a264 100644 --- a/src/Memoization.cpp +++ b/src/Memoization.cpp @@ -470,7 +470,7 @@ Stmt inject_memoization(const Stmt &s, const std::map &en InjectMemoization injector(env, memoize_instance++, name, outputs); - return injector.mutate(s); + return injector(s); } namespace { @@ -563,9 +563,7 @@ class RewriteMemoizedAllocations : public IRMutator { Stmt rewrite_memoized_allocations(const Stmt &s, const std::map &env) { - RewriteMemoizedAllocations rewriter(env); - - return rewriter.mutate(s); + return RewriteMemoizedAllocations(env)(s); } } // namespace Internal diff --git a/src/OffloadGPULoops.cpp b/src/OffloadGPULoops.cpp index 376dcd0b9949..7351a038e019 100644 --- a/src/OffloadGPULoops.cpp +++ b/src/OffloadGPULoops.cpp @@ -44,7 +44,7 @@ class ExtractBounds : public IRVisitor { } } -private: +protected: bool found_shared = false; using IRVisitor::visit; @@ -87,6 +87,7 @@ class ExtractBounds : public IRVisitor { }; class InjectGpuOffload : public IRMutator { +protected: /** Child code generator for device kernels. */ map> cgdev; @@ -131,7 +132,7 @@ class InjectGpuOffload : public IRMutator { << "A concrete device API should have been selected before codegen."; ExtractBounds bounds; - loop->accept(&bounds); + bounds(loop); debug(2) << "Kernel bounds: (" << bounds.num_threads[0] << ", " << bounds.num_threads[1] << ", " diff --git a/src/OptimizeShuffles.cpp b/src/OptimizeShuffles.cpp index 0a88d02f0b60..83672fa59395 100644 --- a/src/OptimizeShuffles.cpp +++ b/src/OptimizeShuffles.cpp @@ -144,7 +144,7 @@ class OptimizeShuffles : public IRMutator { } // namespace Stmt optimize_shuffles(Stmt s, int lut_alignment) { - s = OptimizeShuffles(lut_alignment).mutate(s); + s = OptimizeShuffles(lut_alignment)(s); return s; } diff --git a/src/ParallelRVar.cpp b/src/ParallelRVar.cpp index f2240fb9e036..538cd144f449 100644 --- a/src/ParallelRVar.cpp +++ b/src/ParallelRVar.cpp @@ -110,7 +110,7 @@ bool can_parallelize_rvar(const string &v, // Make an expr representing the store done by a different thread. RenameFreeVars renamer; - auto other_store = renamer.mutate(args); + auto other_store = renamer(args); // Construct an expression which is true when the two threads are // in fact two different threads. We'll use this liberally in the @@ -147,7 +147,7 @@ bool can_parallelize_rvar(const string &v, // Add the definition's predicate if there is any if (pred.defined() || !is_const_one(pred)) { const Expr &this_pred = pred; - Expr other_pred = renamer.mutate(pred); + Expr other_pred = renamer(pred); debug(3) << "......this thread predicate: " << this_pred << "\n"; debug(3) << "......other thread predicate: " << other_pred << "\n"; hazard = hazard && this_pred && other_pred; @@ -156,7 +156,7 @@ bool can_parallelize_rvar(const string &v, debug(3) << "Attempting to falsify: " << hazard << "\n"; // Pull out common non-boolean terms hazard = common_subexpression_elimination(hazard); - hazard = SubstituteInBooleanLets().mutate(hazard); + hazard = SubstituteInBooleanLets()(hazard); hazard = simplify(hazard, bounds); debug(3) << "Simplified to: " << hazard << "\n"; diff --git a/src/PartitionLoops.cpp b/src/PartitionLoops.cpp index 8f80ea42cd85..e31b8b9387b2 100644 --- a/src/PartitionLoops.cpp +++ b/src/PartitionLoops.cpp @@ -1165,16 +1165,16 @@ bool has_likely_tag(const Expr &e, const Scope<> &scope) { } Stmt partition_loops(Stmt s) { - s = LowerLikelyIfInnermost().mutate(s); + s = LowerLikelyIfInnermost()(s); // Walk inwards to the first loop before doing any more work. s = mutate_with(s, [](auto *self, const For *op) { Stmt s = op; - s = MarkClampedRampsAsLikely().mutate(s); - s = ExpandSelects().mutate(s); - s = PartitionLoops().mutate(s); - s = RenormalizeGPULoops().mutate(s); - s = CollapseSelects().mutate(s); + s = MarkClampedRampsAsLikely()(s); + s = ExpandSelects()(s); + s = PartitionLoops()(s); + s = RenormalizeGPULoops()(s); + s = CollapseSelects()(s); return s; }); diff --git a/src/Prefetch.cpp b/src/Prefetch.cpp index c0eedf50b817..6ab41b3deade 100644 --- a/src/Prefetch.cpp +++ b/src/Prefetch.cpp @@ -1,5 +1,6 @@ #include #include +#include #include #include "Bounds.h" @@ -19,6 +20,7 @@ namespace Internal { using std::map; using std::set; using std::string; +using std::unordered_map; using std::vector; /** @@ -42,12 +44,15 @@ namespace { // Collect the bounds of all the externally referenced buffers in a stmt. class CollectExternalBufferBounds : public IRVisitor { public: - map buffers; + unordered_map buffers; using IRVisitor::visit; void add_buffer_bounds(const string &name, const Buffer<> &image, const Parameter ¶m, int dims) { - Box b; + if (buffers.find(name) != buffers.end()) { + return; + } + Box b(dims); for (int i = 0; i < dims; ++i) { string dim_name = std::to_string(i); Expr buf_min_i = Variable::make(Int(32), concat_strings(name, ".min.", i), @@ -55,7 +60,7 @@ class CollectExternalBufferBounds : public IRVisitor { Expr buf_extent_i = Variable::make(Int(32), concat_strings(name, ".extent.", i), image, param, ReductionDomain()); Expr buf_max_i = buf_min_i + buf_extent_i - 1; - b.push_back(Interval(buf_min_i, buf_max_i)); + b[i] = Interval(buf_min_i, buf_max_i); } buffers.emplace(name, b); } @@ -74,13 +79,13 @@ class CollectExternalBufferBounds : public IRVisitor { class InjectPrefetch : public IRMutator { public: - InjectPrefetch(const map &e, const map &buffers) + InjectPrefetch(const map &e, const unordered_map &buffers) : env(e), external_buffers(buffers) { } -private: +protected: const map &env; - const map &external_buffers; + const unordered_map &external_buffers; Scope buffer_bounds; using IRMutator::visit; @@ -187,7 +192,7 @@ class InjectPlaceholderPrefetch : public IRMutator { : env(e), prefix(prefix), prefetch_list(prefetches) { } -private: +protected: const map &env; const string &prefix; const vector &prefetch_list; @@ -260,6 +265,7 @@ class InjectPlaceholderPrefetch : public IRMutator { // Reduce the prefetch dimension if bigger than 'max_dim'. It keeps the 'max_dim' // innermost dimensions and replaces the rests with for-loops. class ReducePrefetchDimension : public IRMutator { +protected: using IRMutator::visit; const size_t max_dim; @@ -321,6 +327,7 @@ class ReducePrefetchDimension : public IRMutator { // prefetch. This will split the prefetch call into multiple calls by adding // an outer for-loop around the prefetch. class SplitPrefetch : public IRMutator { +protected: using IRMutator::visit; Expr max_byte_size; @@ -400,6 +407,7 @@ void traverse_block(const Stmt &s, Fn &&f) { } class HoistPrefetches : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const Block *op) override { @@ -432,14 +440,14 @@ class HoistPrefetches : public IRMutator { Stmt inject_placeholder_prefetch(const Stmt &s, const map &env, const string &prefix, const vector &prefetches) { - Stmt stmt = InjectPlaceholderPrefetch(env, prefix, prefetches).mutate(s); + Stmt stmt = InjectPlaceholderPrefetch(env, prefix, prefetches)(s); return stmt; } Stmt inject_prefetch(const Stmt &s, const map &env) { CollectExternalBufferBounds finder; s.accept(&finder); - return InjectPrefetch(env, finder.buffers).mutate(s); + return InjectPrefetch(env, finder.buffers)(s); } Stmt reduce_prefetch_dimension(Stmt stmt, const Target &t) { @@ -461,17 +469,17 @@ Stmt reduce_prefetch_dimension(Stmt stmt, const Target &t) { } internal_assert(max_dim > 0); - stmt = ReducePrefetchDimension(max_dim).mutate(stmt); + stmt = ReducePrefetchDimension(max_dim)(stmt); if (max_byte_size.defined()) { // If the max byte size is specified, we may need to tile // the prefetch - stmt = SplitPrefetch(max_byte_size).mutate(stmt); + stmt = SplitPrefetch(max_byte_size)(stmt); } return stmt; } Stmt hoist_prefetches(const Stmt &s) { - return HoistPrefetches().mutate(s); + return HoistPrefetches()(s); } } // namespace Internal diff --git a/src/Profiling.cpp b/src/Profiling.cpp index a6bf13483751..2ce8c7c5a194 100644 --- a/src/Profiling.cpp +++ b/src/Profiling.cpp @@ -611,7 +611,7 @@ Stmt inject_profiling(const Stmt &stmt, const string &pipeline_name, const std:: Names names(pipeline_name); InjectProfiling profiling(names, env); - Stmt s = profiling.mutate(stmt); + Stmt s = profiling(stmt); int num_funcs = (int)(profiling.indices.size()); diff --git a/src/PurifyIndexMath.cpp b/src/PurifyIndexMath.cpp index 1ea205a6ff6b..0bdf7f60691d 100644 --- a/src/PurifyIndexMath.cpp +++ b/src/PurifyIndexMath.cpp @@ -1,7 +1,6 @@ #include "PurifyIndexMath.h" #include "IRMutator.h" #include "IROperator.h" -#include "Simplify.h" namespace Halide { namespace Internal { @@ -26,7 +25,7 @@ class PurifyIndexMath : public IRMutator { } // namespace Expr purify_index_math(const Expr &s) { - return PurifyIndexMath().mutate(s); + return PurifyIndexMath()(s); } } // namespace Internal diff --git a/src/Qualify.cpp b/src/Qualify.cpp index 470162909f94..ef4a32abb355 100644 --- a/src/Qualify.cpp +++ b/src/Qualify.cpp @@ -36,8 +36,7 @@ class QualifyExpr : public IRMutator { } // namespace Expr qualify(const string &prefix, const Expr &value) { - QualifyExpr q(prefix); - return q.mutate(value); + return QualifyExpr(prefix)(value); } } // namespace Internal diff --git a/src/Random.cpp b/src/Random.cpp index 57eeb69f9210..e97d43ce0680 100644 --- a/src/Random.cpp +++ b/src/Random.cpp @@ -141,8 +141,7 @@ class LowerRandom : public IRMutator { } // namespace Expr lower_random(const Expr &e, const vector &free_vars, int tag) { - LowerRandom r(free_vars, tag); - return r.mutate(e); + return LowerRandom(free_vars, tag)(e); } } // namespace Internal diff --git a/src/RebaseLoopsToZero.cpp b/src/RebaseLoopsToZero.cpp index 49f97126bb93..dfc0fb3d0709 100644 --- a/src/RebaseLoopsToZero.cpp +++ b/src/RebaseLoopsToZero.cpp @@ -46,7 +46,7 @@ class RebaseLoopsToZero : public IRMutator { } // namespace Stmt rebase_loops_to_zero(const Stmt &s) { - return RebaseLoopsToZero().mutate(s); + return RebaseLoopsToZero()(s); } } // namespace Internal diff --git a/src/Reduction.cpp b/src/Reduction.cpp index bedb51065694..b0dd66ee823f 100644 --- a/src/Reduction.cpp +++ b/src/Reduction.cpp @@ -102,32 +102,32 @@ struct ReductionDomainContents { } // Pass an IRVisitor through to all Exprs referenced in the ReductionDomainContents - void accept(IRVisitor *visitor) { + void accept(IRVisitor &visitor) { for (const ReductionVariable &rvar : domain) { if (rvar.min.defined()) { - rvar.min.accept(visitor); + visitor(rvar.min); } if (rvar.extent.defined()) { - rvar.extent.accept(visitor); + visitor(rvar.extent); } } if (predicate.defined()) { - predicate.accept(visitor); + visitor(predicate); } } // Pass an IRMutator through to all Exprs referenced in the ReductionDomainContents - void mutate(IRMutator *mutator) { + void mutate(IRMutator &mutator) { for (ReductionVariable &rvar : domain) { if (rvar.min.defined()) { - rvar.min = mutator->mutate(rvar.min); + rvar.min = mutator(rvar.min); } if (rvar.extent.defined()) { - rvar.extent = mutator->mutate(rvar.extent); + rvar.extent = mutator(rvar.extent); } } if (predicate.defined()) { - predicate = mutator->mutate(predicate); + predicate = mutator(predicate); } } }; @@ -196,7 +196,7 @@ class DropSelfReferences : public IRMutator { void ReductionDomain::set_predicate(const Expr &p) { // The predicate can refer back to the RDom. We need to break // those cycles to prevent a leak. - contents->predicate = DropSelfReferences(p, *this).mutate(p); + contents->predicate = DropSelfReferences(p, *this)(p); } void ReductionDomain::where(Expr predicate) { @@ -223,13 +223,13 @@ bool ReductionDomain::frozen() const { void ReductionDomain::accept(IRVisitor *visitor) const { if (contents.defined()) { - contents->accept(visitor); + contents->accept(*visitor); } } void ReductionDomain::mutate(IRMutator *mutator) { if (contents.defined()) { - contents->mutate(mutator); + contents->mutate(*mutator); } } diff --git a/src/RemoveDeadAllocations.cpp b/src/RemoveDeadAllocations.cpp index 33a1a0190b07..74a5c32751d2 100644 --- a/src/RemoveDeadAllocations.cpp +++ b/src/RemoveDeadAllocations.cpp @@ -9,6 +9,7 @@ namespace Internal { namespace { class RemoveDeadAllocations : public IRMutator { +protected: using IRMutator::visit; Scope allocs; @@ -88,7 +89,7 @@ class RemoveDeadAllocations : public IRMutator { } // namespace Stmt remove_dead_allocations(const Stmt &s) { - return RemoveDeadAllocations().mutate(s); + return RemoveDeadAllocations()(s); } } // namespace Internal diff --git a/src/RemoveExternLoops.cpp b/src/RemoveExternLoops.cpp index 9fb0e187b3eb..c3aab4db7c75 100644 --- a/src/RemoveExternLoops.cpp +++ b/src/RemoveExternLoops.cpp @@ -22,7 +22,7 @@ class RemoveExternLoops : public IRMutator { } // namespace Stmt remove_extern_loops(const Stmt &s) { - return RemoveExternLoops().mutate(s); + return RemoveExternLoops()(s); } } // namespace Internal diff --git a/src/RemoveUndef.cpp b/src/RemoveUndef.cpp index 9667aafe891a..512c984427a7 100644 --- a/src/RemoveUndef.cpp +++ b/src/RemoveUndef.cpp @@ -628,7 +628,7 @@ class RemoveUndef : public IRMutator { Stmt remove_undef(Stmt s) { RemoveUndef r; - s = r.mutate(s); + s = r(s); internal_assert(!r.predicate.defined()) << "Undefined expression leaked outside of a Store node: " << r.predicate << "\n"; diff --git a/src/Schedule.cpp b/src/Schedule.cpp index 9f5f51d7a043..a2583d0fb732 100644 --- a/src/Schedule.cpp +++ b/src/Schedule.cpp @@ -250,33 +250,33 @@ struct FuncScheduleContents { } // Pass an IRMutator through to all Exprs referenced in the FuncScheduleContents - void mutate(IRMutator *mutator) { + void mutate(IRMutator &mutator) { for (Bound &b : bounds) { if (b.min.defined()) { - b.min = mutator->mutate(b.min); + b.min = mutator(b.min); } if (b.extent.defined()) { - b.extent = mutator->mutate(b.extent); + b.extent = mutator(b.extent); } if (b.modulus.defined()) { - b.modulus = mutator->mutate(b.modulus); + b.modulus = mutator(b.modulus); } if (b.remainder.defined()) { - b.remainder = mutator->mutate(b.remainder); + b.remainder = mutator(b.remainder); } } for (Bound &b : estimates) { if (b.min.defined()) { - b.min = mutator->mutate(b.min); + b.min = mutator(b.min); } if (b.extent.defined()) { - b.extent = mutator->mutate(b.extent); + b.extent = mutator(b.extent); } if (b.modulus.defined()) { - b.modulus = mutator->mutate(b.modulus); + b.modulus = mutator(b.modulus); } if (b.remainder.defined()) { - b.remainder = mutator->mutate(b.remainder); + b.remainder = mutator(b.remainder); } } } @@ -313,23 +313,23 @@ struct StageScheduleContents { } // Pass an IRMutator through to all Exprs referenced in the StageScheduleContents - void mutate(IRMutator *mutator) { + void mutate(IRMutator &mutator) { for (ReductionVariable &r : rvars) { if (r.min.defined()) { - r.min = mutator->mutate(r.min); + r.min = mutator(r.min); } if (r.extent.defined()) { - r.extent = mutator->mutate(r.extent); + r.extent = mutator(r.extent); } } for (Split &s : splits) { if (s.factor.defined()) { - s.factor = mutator->mutate(s.factor); + s.factor = mutator(s.factor); } } for (PrefetchDirective &p : prefetches) { if (p.offset.defined()) { - p.offset = mutator->mutate(p.offset); + p.offset = mutator(p.offset); } } } @@ -521,7 +521,7 @@ void FuncSchedule::accept(IRVisitor *visitor) const { void FuncSchedule::mutate(IRMutator *mutator) { if (contents.defined()) { - contents->mutate(mutator); + contents->mutate(*mutator); } } @@ -665,7 +665,7 @@ void StageSchedule::accept(IRVisitor *visitor) const { void StageSchedule::mutate(IRMutator *mutator) { if (contents.defined()) { - contents->mutate(mutator); + contents->mutate(*mutator); } } diff --git a/src/ScheduleFunctions.cpp b/src/ScheduleFunctions.cpp index 1a9e0858c4d6..a0b4a4c18876 100644 --- a/src/ScheduleFunctions.cpp +++ b/src/ScheduleFunctions.cpp @@ -133,7 +133,7 @@ class SubstituteIn : public IRGraphMutator { }; Stmt substitute_in(const string &name, const Expr &value, bool calls, bool provides, const Stmt &s) { - return SubstituteIn(name, value, calls, provides).mutate(s); + return SubstituteIn(name, value, calls, provides)(s); } class AddPredicates : public IRGraphMutator { @@ -177,7 +177,7 @@ class AddPredicates : public IRGraphMutator { }; Stmt add_predicates(const Expr &cond, const Function &func, ApplySplitResult::Type type, const Stmt &s) { - return AddPredicates(cond, func, type).mutate(s); + return AddPredicates(cond, func, type)(s); } // Build a loop nest about a provide node using a schedule @@ -1004,7 +1004,7 @@ Stmt inject_stmt(Stmt root, Stmt injected, const LoopLevel &level) { return Block::make(root, injected); } InjectStmt injector(injected, level); - root = injector.mutate(root); + root = injector(root); internal_assert(injector.found_level); return root; } @@ -1091,7 +1091,7 @@ Stmt substitute_fused_bounds(Stmt s, const map &replacements) } } subs(replacements); - return subs.mutate(s); + return subs(s); } // Add letstmts inside each parent loop that define the corresponding child loop @@ -1128,7 +1128,7 @@ Stmt add_loop_var_aliases(Stmt s, const map> &loop_var_alias } } add_aliases(loop_var_aliases); - return add_aliases.mutate(s); + return add_aliases(s); } // Shift the iteration domain of a loop nest by some factor. @@ -1161,8 +1161,7 @@ class ShiftLoopNest : public IRMutator { if (shifts.empty()) { return node; } - ShiftLoopNest visitor(shifts); - return visitor.mutate(node); + return ShiftLoopNest(shifts)(node); } }; @@ -2612,7 +2611,7 @@ Stmt schedule_functions(const vector &outputs, } else { debug(1) << "Injecting realization of " << funcs << "\n"; InjectFunctionRealization injector(funcs, is_output_list, target, env); - s = injector.mutate(s); + s = injector(s); internal_assert(injector.found_store_level() && injector.found_compute_level() && injector.found_hoist_storage_level()); } @@ -2625,7 +2624,7 @@ Stmt schedule_functions(const vector &outputs, s = root_loop->body; // We can also remove all the loops over __outermost now. - s = RemoveLoopsOverOutermost().mutate(s); + s = RemoveLoopsOverOutermost()(s); return s; } diff --git a/src/SelectGPUAPI.cpp b/src/SelectGPUAPI.cpp index ec73c883e955..4c764f6bc619 100644 --- a/src/SelectGPUAPI.cpp +++ b/src/SelectGPUAPI.cpp @@ -50,7 +50,7 @@ class SelectGPUAPI : public IRMutator { } // namespace Stmt select_gpu_api(const Stmt &s, const Target &t) { - return SelectGPUAPI(t).mutate(s); + return SelectGPUAPI(t)(s); } } // namespace Internal diff --git a/src/Simplify.cpp b/src/Simplify.cpp index edaccbcefc16..a7d7d7f6faa0 100644 --- a/src/Simplify.cpp +++ b/src/Simplify.cpp @@ -16,7 +16,6 @@ using std::string; using std::vector; Simplify::Simplify(const Scope *bi, const Scope *ai) { - // Only respect the constant bounds from the containing scope. for (auto iter = bi->cbegin(); iter != bi->cend(); ++iter) { ExprInfo info; @@ -450,7 +449,7 @@ bool can_prove(Expr e, const Scope &bounds) { std::vector> out_vars; } renamer; - e = renamer.mutate(e); + e = renamer(e); // Look for a concrete counter-example with random probing static std::mt19937 rng(0); diff --git a/src/SimplifyCorrelatedDifferences.cpp b/src/SimplifyCorrelatedDifferences.cpp index 3afe5d84dcce..b3c016ea3c57 100644 --- a/src/SimplifyCorrelatedDifferences.cpp +++ b/src/SimplifyCorrelatedDifferences.cpp @@ -20,6 +20,7 @@ using std::string; using std::vector; class PartiallyCancelDifferences : public IRMutator { +protected: using IRMutator::visit; // Symbols used by rewrite rules @@ -65,6 +66,7 @@ class PartiallyCancelDifferences : public IRMutator { }; class SimplifyCorrelatedDifferences : public IRMutator { +protected: using IRMutator::visit; string loop_var; @@ -177,6 +179,7 @@ class SimplifyCorrelatedDifferences : public IRMutator { // Add the names of any free variables in an expr to the provided set void track_free_vars(const Expr &e, std::set *vars) { class TrackFreeVars : public IRVisitor { + protected: using IRVisitor::visit; void visit(const Variable *op) override { if (!scope.contains(op->name)) { @@ -195,7 +198,7 @@ class SimplifyCorrelatedDifferences : public IRMutator { : vars(vars) { } } tracker(vars); - e.accept(&tracker); + tracker(e); } Expr cancel_correlated_subexpression(Expr e, const Expr &a, const Expr &b, bool correlated) { @@ -224,7 +227,7 @@ class SimplifyCorrelatedDifferences : public IRMutator { } e = common_subexpression_elimination(e); e = solve_expression(e, loop_var).result; - e = PartiallyCancelDifferences().mutate(e); + e = PartiallyCancelDifferences()(e); e = simplify(e); const bool check_non_monotonic = debug_is_active(1) || get_compiler_logger() != nullptr; @@ -308,11 +311,11 @@ class SimplifyCorrelatedDifferences : public IRMutator { } // namespace Stmt simplify_correlated_differences(const Stmt &stmt) { - return SimplifyCorrelatedDifferences().mutate(stmt); + return SimplifyCorrelatedDifferences()(stmt); } Expr bound_correlated_differences(const Expr &expr) { - return PartiallyCancelDifferences().mutate(expr); + return PartiallyCancelDifferences()(expr); } } // namespace Internal diff --git a/src/SkipStages.cpp b/src/SkipStages.cpp index 7fa3e05dd99b..6f59ca538216 100644 --- a/src/SkipStages.cpp +++ b/src/SkipStages.cpp @@ -835,7 +835,7 @@ Stmt skip_stages(const Stmt &stmt, } SkipStages skipper(analysis, name_for_id); - stmt = skipper.mutate(stmt); + stmt = skipper(stmt); stmt = skipper.emit_outermost_defs(stmt); return stmt; }; diff --git a/src/SlidingWindow.cpp b/src/SlidingWindow.cpp index a47ecb6b1014..f9122f572496 100644 --- a/src/SlidingWindow.cpp +++ b/src/SlidingWindow.cpp @@ -86,7 +86,7 @@ class ExpandExpr : public IRMutator { // Perform all the substitutions in a scope Expr expand_expr(const Expr &e, const Scope &scope) { ExpandExpr ee(scope); - Expr result = ee.mutate(e); + Expr result = ee(e); debug(4) << "Expanded " << e << " into " << result << "\n"; return result; } @@ -614,7 +614,7 @@ class SlidingWindowOnFunctionAndLoop : public IRMutator { Interval new_bounds; Stmt translate_loop(const Stmt &s) { - return RollFunc(func, dim_idx, loop_var, old_bounds, new_bounds).mutate(s); + return RollFunc(func, dim_idx, loop_var, old_bounds, new_bounds)(s); } }; @@ -819,7 +819,7 @@ class SlidingWindow : public IRMutator { set &slid_dims = slid_dimensions[func.name()]; size_t old_slid_dims_size = slid_dims.size(); SlidingWindowOnFunctionAndLoop slider(func, name, prev_loop_min, slid_dims); - body = slider.mutate(body); + body = slider(body); if (func.schedule().memory_type() == MemoryType::Register && slider.old_bounds.has_lower_bound()) { @@ -847,7 +847,7 @@ class SlidingWindow : public IRMutator { {name + ".loop_min", loop_min}, }, body); - body = SubstitutePrefetchVar(name, new_name).mutate(body); + body = SubstitutePrefetchVar(name, new_name)(body); name = new_name; @@ -923,7 +923,7 @@ class AddLoopMinOrig : public IRMutator { } // namespace Stmt sliding_window(const Stmt &s, const map &env) { - return SlidingWindow(env).mutate(AddLoopMinOrig().mutate(s)); + return SlidingWindow(env)(AddLoopMinOrig()(s)); } } // namespace Internal diff --git a/src/Solve.cpp b/src/Solve.cpp index 8245f980ee46..a0b91be7a287 100644 --- a/src/Solve.cpp +++ b/src/Solve.cpp @@ -61,7 +61,7 @@ class SolveExpression : public IRMutator { // Has the solve failed. bool failed = false; -private: +protected: // The variable we're solving for. string var; @@ -1136,7 +1136,7 @@ class SolveForInterval : public IRVisitor { SolverResult solve_expression(const Expr &e, const std::string &variable, const Scope &scope) { SolveExpression solver(variable, scope); - Expr new_e = solver.mutate(e); + Expr new_e = solver(e); // The process has expanded lets. Re-collect them. new_e = common_subexpression_elimination(new_e); debug(3) << "Solved expr for " << variable << " :\n" diff --git a/src/SplitTuples.cpp b/src/SplitTuples.cpp index 314780d981bd..99f056085cae 100644 --- a/src/SplitTuples.cpp +++ b/src/SplitTuples.cpp @@ -410,7 +410,7 @@ class SplitScatterGather : public IRMutator { vector vars; for (extractor.idx = 0; extractor.idx < size; extractor.idx++) { string name = unique_name(op->name + "." + std::to_string(extractor.idx)); - lets.emplace_back(name, extractor.mutate(op->value)); + lets.emplace_back(name, extractor(op->value)); vars.push_back(Variable::make(op->value.type(), name)); } @@ -477,15 +477,15 @@ class SplitScatterGather : public IRMutator { vector args = op->args; for (Expr &a : args) { string name = unique_name('t'); - exprs.push_back(extractor.mutate(a)); + exprs.push_back(extractor(a)); names.push_back(name); a = Variable::make(a.type(), name); } vector values = op->values; for (Expr &v : values) { - v = extractor.mutate(v); + v = extractor(v); string name = unique_name('t'); - exprs.push_back(extractor.mutate(v)); + exprs.push_back(extractor(v)); names.push_back(name); v = Variable::make(v.type(), name); } @@ -526,8 +526,8 @@ class SplitScatterGather : public IRMutator { } // namespace Stmt split_tuples(const Stmt &stmt, const map &env) { - Stmt s = SplitTuples(env).mutate(stmt); - s = SplitScatterGather().mutate(s); + Stmt s = SplitTuples(env)(stmt); + s = SplitScatterGather()(s); return s; } diff --git a/src/StorageFlattening.cpp b/src/StorageFlattening.cpp index 987d22220791..da14d1b52530 100644 --- a/src/StorageFlattening.cpp +++ b/src/StorageFlattening.cpp @@ -414,7 +414,7 @@ class FlattenDimensions : public IRMutator { }; class HoistStorage : public IRMutator { - +protected: struct HoistedAllocationInfo { string name; Type type; @@ -580,6 +580,7 @@ class HoistStorage : public IRMutator { // Realizations, stores, and loads must all be on types that are // multiples of 8-bits. This really only affects bools class PromoteToMemoryType : public IRMutator { +protected: using IRMutator::visit; Type upgrade(Type t) { @@ -640,9 +641,9 @@ Stmt storage_flattening(Stmt s, tuple_env[p.first] = {p.second, 0}; } } - s = FlattenDimensions(tuple_env, outputs, target).mutate(s); - s = HoistStorage().mutate(s); - s = PromoteToMemoryType().mutate(s); + s = FlattenDimensions(tuple_env, outputs, target)(s); + s = HoistStorage()(s); + s = PromoteToMemoryType()(s); return s; } diff --git a/src/StorageFolding.cpp b/src/StorageFolding.cpp index e97b06a8a6b9..4b4fa8152d4a 100644 --- a/src/StorageFolding.cpp +++ b/src/StorageFolding.cpp @@ -691,8 +691,7 @@ class AttemptStorageFoldingOfFunction : public IRMutator { op->name, sema_var, dim, - storage_dim) - .mutate(body); + storage_dim)(body); if (storage_dim.fold_forward) { can_fold_forwards = true; @@ -791,7 +790,7 @@ class AttemptStorageFoldingOfFunction : public IRMutator { } else { head = dynamic_footprint; } - body = FoldStorageOfFunction(func.name(), (int)i - 1, factor, head).mutate(body); + body = FoldStorageOfFunction(func.name(), (int)i - 1, factor, head)(body); } // If the producer is async, it can run ahead by @@ -961,7 +960,7 @@ class StorageFolding : public IRMutator { } else { debug(3) << "Attempting to fold " << op->name << " automatically or explicitly\n"; } - body = folder.mutate(body); + body = folder(body); if (body.same_as(op->body)) { return op; @@ -1034,8 +1033,8 @@ class RemoveSlidingWindowMarkers : public IRMutator { } // namespace Stmt storage_folding(const Stmt &s, const std::map &env) { - Stmt stmt = StorageFolding(env).mutate(s); - stmt = RemoveSlidingWindowMarkers().mutate(stmt); + Stmt stmt = StorageFolding(env)(s); + stmt = RemoveSlidingWindowMarkers()(stmt); return stmt; } diff --git a/src/StrictifyFloat.cpp b/src/StrictifyFloat.cpp index 37263d00c89b..9deb86679808 100644 --- a/src/StrictifyFloat.cpp +++ b/src/StrictifyFloat.cpp @@ -127,7 +127,7 @@ class AnyStrictIntrinsics : public IRVisitor { } // namespace Expr strictify_float(const Expr &e) { - return Strictify{}.mutate(e); + return Strictify{}(e); } Expr unstrictify_float(const Call *op) { diff --git a/src/StripAsserts.cpp b/src/StripAsserts.cpp index e8f101fcfc0c..f1ffb8967ca9 100644 --- a/src/StripAsserts.cpp +++ b/src/StripAsserts.cpp @@ -109,7 +109,7 @@ class StripAsserts : public IRMutator { } // namespace Stmt strip_asserts(const Stmt &s) { - return StripAsserts().mutate(s); + return StripAsserts()(s); } } // namespace Internal diff --git a/src/Substitute.cpp b/src/Substitute.cpp index 9b280b0b0483..3781914944cf 100644 --- a/src/Substitute.cpp +++ b/src/Substitute.cpp @@ -12,6 +12,7 @@ using std::string; namespace { class Substitute : public IRMutator { +protected: const map &replace; Scope<> hidden; @@ -104,24 +105,24 @@ Expr substitute(const string &name, const Expr &replacement, const Expr &expr) { map m; m[name] = replacement; Substitute s(m); - return s.mutate(expr); + return s(expr); } Stmt substitute(const string &name, const Expr &replacement, const Stmt &stmt) { map m; m[name] = replacement; Substitute s(m); - return s.mutate(stmt); + return s(stmt); } Expr substitute(const map &m, const Expr &expr) { Substitute s(m); - return s.mutate(expr); + return s(expr); } Stmt substitute(const map &m, const Stmt &stmt) { Substitute s(m); - return s.mutate(stmt); + return s(stmt); } namespace { @@ -150,6 +151,7 @@ namespace { /** Substitute an expr for a var in a graph. */ class GraphSubstitute : public IRGraphMutator { +protected: string var; Expr value; @@ -202,25 +204,25 @@ class GraphSubstituteExpr : public IRGraphMutator { } // namespace Expr graph_substitute(const string &name, const Expr &replacement, const Expr &expr) { - return GraphSubstitute(name, replacement).mutate(expr); + return GraphSubstitute(name, replacement)(expr); } Stmt graph_substitute(const string &name, const Expr &replacement, const Stmt &stmt) { - return GraphSubstitute(name, replacement).mutate(stmt); + return GraphSubstitute(name, replacement)(stmt); } Expr graph_substitute(const Expr &find, const Expr &replacement, const Expr &expr) { - return GraphSubstituteExpr(find, replacement).mutate(expr); + return GraphSubstituteExpr(find, replacement)(expr); } Stmt graph_substitute(const Expr &find, const Expr &replacement, const Stmt &stmt) { - return GraphSubstituteExpr(find, replacement).mutate(stmt); + return GraphSubstituteExpr(find, replacement)(stmt); } namespace { class SubstituteInAllLets : public IRGraphMutator { - +protected: using IRGraphMutator::visit; Expr visit(const Let *op) override { @@ -233,11 +235,11 @@ class SubstituteInAllLets : public IRGraphMutator { } // namespace Expr substitute_in_all_lets(const Expr &expr) { - return SubstituteInAllLets().mutate(expr); + return SubstituteInAllLets()(expr); } Stmt substitute_in_all_lets(const Stmt &stmt) { - return SubstituteInAllLets().mutate(stmt); + return SubstituteInAllLets()(stmt); } } // namespace Internal diff --git a/src/Tracing.cpp b/src/Tracing.cpp index 0bc9086d8635..59c310f757f2 100644 --- a/src/Tracing.cpp +++ b/src/Tracing.cpp @@ -350,7 +350,7 @@ Stmt inject_tracing(Stmt s, const string &pipeline_name, bool trace_pipeline, } // Inject tracing calls - s = tracing.mutate(s); + s = tracing(s); // Strip off the dummy realize blocks s = mutate_with(s, [&](auto *self, const Realize *op) { diff --git a/src/TrimNoOps.cpp b/src/TrimNoOps.cpp index 1842a702fab4..13d358a4f0bd 100644 --- a/src/TrimNoOps.cpp +++ b/src/TrimNoOps.cpp @@ -105,7 +105,7 @@ class IsNoOp : public IRVisitor { Expr equivalent_load = Load::make(op->value.type(), op->name, op->index, Buffer<>(), Parameter(), op->predicate, op->alignment); Expr is_no_op = equivalent_load == op->value; - is_no_op = StripIdentities().mutate(is_no_op); + is_no_op = StripIdentities()(is_no_op); // We need to call CSE since sometimes we have "let" stmt on the RHS // that makes the expr harder to solve, i.e. the solver will just give up // and return a conservative false on call to and_condition_over_domain(). @@ -405,7 +405,7 @@ class TrimNoOps : public IRMutator { // Simplify the body to take advantage of the fact that the // loop range is now truncated - body = simplify(SimplifyUsingBounds(op->name, i).mutate(body)); + body = simplify(SimplifyUsingBounds(op->name, i)(body)); string new_min_name = unique_name(op->name + ".new_min"); string new_max_name = unique_name(op->name + ".new_max"); @@ -445,9 +445,8 @@ class TrimNoOps : public IRMutator { } // namespace -Stmt trim_no_ops(Stmt s) { - s = TrimNoOps().mutate(s); - return s; +Stmt trim_no_ops(const Stmt &s) { + return TrimNoOps()(s); } } // namespace Internal diff --git a/src/TrimNoOps.h b/src/TrimNoOps.h index 51d264cd03fb..548e2d383857 100644 --- a/src/TrimNoOps.h +++ b/src/TrimNoOps.h @@ -13,7 +13,7 @@ namespace Internal { /** Truncate loop bounds to the region over which they actually do * something. For examples see test/correctness/trim_no_ops.cpp */ -Stmt trim_no_ops(Stmt s); +Stmt trim_no_ops(const Stmt &s); } // namespace Internal } // namespace Halide diff --git a/src/UniquifyVariableNames.cpp b/src/UniquifyVariableNames.cpp index 91f0279de04c..6eec453afc63 100644 --- a/src/UniquifyVariableNames.cpp +++ b/src/UniquifyVariableNames.cpp @@ -5,7 +5,6 @@ #include "IRVisitor.h" #include "Scope.h" #include "Var.h" -#include namespace Halide { namespace Internal { @@ -16,7 +15,7 @@ using std::vector; namespace { class UniquifyVariableNames : public IRMutator { - +protected: using IRMutator::visit; // The mapping from old names to new names @@ -119,7 +118,7 @@ class UniquifyVariableNames : public IRMutator { }; class FindFreeVars : public IRVisitor { - +protected: using IRVisitor::visit; Scope<> scope; @@ -168,8 +167,7 @@ class FindFreeVars : public IRVisitor { Stmt uniquify_variable_names(const Stmt &s) { FindFreeVars finder; s.accept(&finder); - UniquifyVariableNames u(&finder.free_vars); - return u.mutate(s); + return UniquifyVariableNames(&finder.free_vars)(s); } namespace { diff --git a/src/UnrollLoops.cpp b/src/UnrollLoops.cpp index ffcba564966a..127dac785dea 100644 --- a/src/UnrollLoops.cpp +++ b/src/UnrollLoops.cpp @@ -11,6 +11,7 @@ namespace Internal { namespace { class UnrollLoops : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const For *for_loop) override { @@ -19,8 +20,7 @@ class UnrollLoops : public IRMutator { Expr extent = simplify(for_loop->extent()); const IntImm *e = extent.as(); - internal_assert(e) - << "Loop over " << for_loop->name << " should have had a constant extent\n"; + internal_assert(e) << "Loop over " << for_loop->name << " should have had a constant extent\n"; body = mutate(body); if (e->value == 1) { @@ -53,7 +53,7 @@ class UnrollLoops : public IRMutator { } // namespace Stmt unroll_loops(const Stmt &s) { - Stmt stmt = UnrollLoops().mutate(s); + Stmt stmt = UnrollLoops()(s); // Unrolling duplicates variable names. Other passes assume variable names are unique. return uniquify_variable_names(stmt); } diff --git a/src/UnsafePromises.cpp b/src/UnsafePromises.cpp index c1fdc51d8758..02912594d49b 100644 --- a/src/UnsafePromises.cpp +++ b/src/UnsafePromises.cpp @@ -57,11 +57,11 @@ class LowerSafePromises : public IRMutator { } // namespace Stmt lower_unsafe_promises(const Stmt &s, const Target &t) { - return LowerUnsafePromises(t.has_feature(Target::CheckUnsafePromises)).mutate(s); + return LowerUnsafePromises(t.has_feature(Target::CheckUnsafePromises))(s); } Stmt lower_safe_promises(const Stmt &s) { - return LowerSafePromises().mutate(s); + return LowerSafePromises()(s); } } // namespace Internal diff --git a/src/VectorizeLoops.cpp b/src/VectorizeLoops.cpp index 18243503372b..ebfd63e860bd 100644 --- a/src/VectorizeLoops.cpp +++ b/src/VectorizeLoops.cpp @@ -309,6 +309,7 @@ bool is_interleaved_ramp(const Expr &e, const Scope &scope, InterleavedRam // vector lane. This means loads and stores to them need to be // rewritten slightly. class RewriteAccessToVectorAlloc : public IRMutator { +protected: Expr var; string alloc; int lanes; @@ -363,6 +364,7 @@ class SerializeLoops : public IRMutator { // Wrap a vectorized predicate around a Load/Store node. class PredicateLoadStore : public IRMutator { +protected: string var; Expr vector_predicate; int lanes; @@ -480,6 +482,7 @@ struct VectorizedVar { // Substitutes a vector for a scalar var in a Stmt. Used on the // body of every vectorized loop. class VectorSubs : public IRMutator { +protected: // A list of vectorized loop vars encountered so far. The last // element corresponds to the most inner vectorized loop. std::vector vectorized_vars; @@ -862,12 +865,12 @@ class VectorSubs : public IRMutator { Stmt predicated_stmt; if (vectorize_predicate) { PredicateLoadStore p(vectorized_vars.front().name, cond); - predicated_stmt = p.mutate(then_case); + predicated_stmt = p(then_case); vectorize_predicate = p.is_vectorized(); } if (vectorize_predicate && else_case.defined()) { PredicateLoadStore p(vectorized_vars.front().name, !cond); - predicated_stmt = Block::make(predicated_stmt, p.mutate(else_case)); + predicated_stmt = Block::make(predicated_stmt, p(else_case)); vectorize_predicate = p.is_vectorized(); } @@ -1073,7 +1076,7 @@ class VectorSubs : public IRMutator { // Rewrite loads and stores to this allocation like so: // foo[x] -> foo[x*lanes + v] for (const auto &vv : vectorized_vars) { - body = RewriteAccessToVectorAlloc(vv.name + ".from_zero", op->name, vv.lanes).mutate(body); + body = RewriteAccessToVectorAlloc(vv.name + ".from_zero", op->name, vv.lanes)(body); } body = mutate(body); @@ -1315,7 +1318,7 @@ class VectorSubs : public IRMutator { // better luck vectorizing it. if (serialize_inner_loops) { - s = SerializeLoops().mutate(s); + s = SerializeLoops()(s); } // We'll need the original scalar versions of any containing lets. for (const auto &[var, value] : reverse_view(containing_lets)) { @@ -1411,6 +1414,7 @@ class VectorSubs : public IRMutator { }; // namespace class FindVectorizableExprsInAtomicNode : public IRMutator { +protected: // An Atomic node protects all accesses to a given buffer. We // consider a name "poisoned" if it depends on an access to this // buffer. We can't lift or vectorize anything that has been @@ -1502,6 +1506,7 @@ class FindVectorizableExprsInAtomicNode : public IRMutator { }; class LiftVectorizableExprsOutOfSingleAtomicNode : public IRMutator { +protected: const std::set &liftable; using IRMutator::visit; @@ -1556,6 +1561,7 @@ class LiftVectorizableExprsOutOfSingleAtomicNode : public IRMutator { }; class LiftVectorizableExprsOutOfAllAtomicNodes : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const Atomic *op) override { @@ -1582,6 +1588,7 @@ class LiftVectorizableExprsOutOfAllAtomicNodes : public IRMutator { // Vectorize all loops marked as such in a Stmt class VectorizeLoops : public IRMutator { +protected: using IRMutator::visit; Stmt visit(const For *for_loop) override { @@ -1597,7 +1604,7 @@ class VectorizeLoops : public IRMutator { } VectorizedVar vectorized_var = {for_loop->name, for_loop->min, (int)extent->value}; - stmt = VectorSubs(vectorized_var).mutate(for_loop->body); + stmt = VectorSubs(vectorized_var)(for_loop->body); } else { stmt = IRMutator::visit(for_loop); } @@ -1630,6 +1637,7 @@ bool all_stores_in_scope(const Stmt &stmt, const Scope<> &scope) { /** Drop any atomic nodes protecting buffers that are only accessed * from a single thread. */ class RemoveUnnecessaryAtomics : public IRMutator { +protected: using IRMutator::visit; // Allocations made from within this same thread @@ -1664,7 +1672,7 @@ class RemoveUnnecessaryAtomics : public IRMutator { }; Stmt vectorize_statement(const Stmt &stmt) { - return VectorizeLoops().mutate(stmt); + return VectorizeLoops()(stmt); } } // namespace @@ -1672,9 +1680,9 @@ Stmt vectorize_loops(const Stmt &stmt, const map &env) { // Limit the scope of atomic nodes to just the necessary stuff. // TODO: Should this be an earlier pass? It's probably a good idea // for non-vectorizing stuff too. - Stmt s = LiftVectorizableExprsOutOfAllAtomicNodes(env).mutate(stmt); + Stmt s = LiftVectorizableExprsOutOfAllAtomicNodes(env)(stmt); s = vectorize_statement(s); - s = RemoveUnnecessaryAtomics().mutate(s); + s = RemoveUnnecessaryAtomics()(s); return s; } diff --git a/src/autoschedulers/adams2019/FunctionDAG.cpp b/src/autoschedulers/adams2019/FunctionDAG.cpp index 35b872477150..974182970f54 100644 --- a/src/autoschedulers/adams2019/FunctionDAG.cpp +++ b/src/autoschedulers/adams2019/FunctionDAG.cpp @@ -661,8 +661,8 @@ FunctionDAG::FunctionDAG(const vector &outputs, const Target &target) stage_scope_with_concrete_rvar_bounds.set_containing_scope(&scope); stage_scope_with_symbolic_rvar_bounds.set_containing_scope(&scope); for (const auto &rv : sched.rvars()) { - Expr min = simplify(apply_param_estimates.mutate(rv.min)); - Expr max = simplify(apply_param_estimates.mutate(rv.min + rv.extent - 1)); + Expr min = simplify(apply_param_estimates(rv.min)); + Expr max = simplify(apply_param_estimates(rv.min + rv.extent - 1)); stage_scope_with_concrete_rvar_bounds.push(rv.var, Interval(min, max)); min = Variable::make(Int(32), consumer.name() + "." + rv.var + ".min"); max = Variable::make(Int(32), consumer.name() + "." + rv.var + ".max"); @@ -695,8 +695,8 @@ FunctionDAG::FunctionDAG(const vector &outputs, const Target &target) const auto &req = node.region_required[j]; auto &comp = node.region_computed[j]; comp.depends_on_estimate = depends_on_estimate(comp.in.min) || depends_on_estimate(comp.in.max); - comp.in.min = simplify(apply_param_estimates.mutate(comp.in.min)); - comp.in.max = simplify(apply_param_estimates.mutate(comp.in.max)); + comp.in.min = simplify(apply_param_estimates(comp.in.min)); + comp.in.max = simplify(apply_param_estimates(comp.in.max)); if (equal(comp.in.min, req.min) && equal(comp.in.max, req.max)) { comp.equals_required = true; } else { @@ -921,7 +921,7 @@ FunctionDAG::FunctionDAG(const vector &outputs, const Target &target) stage.index = s; - exprs = apply_param_estimates.mutate(exprs); + exprs = apply_param_estimates(exprs); // For this stage scope we want symbolic bounds for the rvars @@ -949,8 +949,8 @@ FunctionDAG::FunctionDAG(const vector &outputs, const Target &target) << edge.producer->func.name() << " -> " << edge.consumer->name << "\n"; bool min_dependent = depends_on_estimate(in.min); bool max_dependent = depends_on_estimate(in.max); - Expr min_value = simplify(apply_param_estimates.mutate(in.min)); - Expr max_value = simplify(apply_param_estimates.mutate(in.max)); + Expr min_value = simplify(apply_param_estimates(in.min)); + Expr max_value = simplify(apply_param_estimates(in.max)); Edge::BoundInfo min(min_value, *edge.consumer, min_dependent); Edge::BoundInfo max(max_value, *edge.consumer, max_dependent); edge.bounds.emplace_back(std::move(min), std::move(max)); diff --git a/src/autoschedulers/anderson2021/FunctionDAG.cpp b/src/autoschedulers/anderson2021/FunctionDAG.cpp index e127a02a7bd3..07e63280ff12 100644 --- a/src/autoschedulers/anderson2021/FunctionDAG.cpp +++ b/src/autoschedulers/anderson2021/FunctionDAG.cpp @@ -653,8 +653,8 @@ FunctionDAG::FunctionDAG(const vector &outputs, const Target &target) stage_scope_with_concrete_rvar_bounds.set_containing_scope(&scope); stage_scope_with_symbolic_rvar_bounds.set_containing_scope(&scope); for (const auto &rv : sched.rvars()) { - Expr min = simplify(apply_param_estimates.mutate(rv.min)); - Expr max = simplify(apply_param_estimates.mutate(rv.min + rv.extent - 1)); + Expr min = simplify(apply_param_estimates(rv.min)); + Expr max = simplify(apply_param_estimates(rv.min + rv.extent - 1)); stage_scope_with_concrete_rvar_bounds.push(rv.var, Interval(min, max)); min = Variable::make(Int(32), consumer.name() + "." + rv.var + ".min"); max = Variable::make(Int(32), consumer.name() + "." + rv.var + ".max"); @@ -686,8 +686,8 @@ FunctionDAG::FunctionDAG(const vector &outputs, const Target &target) for (int j = 0; j < consumer.dimensions(); j++) { const auto &req = node.region_required[j]; auto &comp = node.region_computed[j]; - comp.in.min = simplify(apply_param_estimates.mutate(comp.in.min)); - comp.in.max = simplify(apply_param_estimates.mutate(comp.in.max)); + comp.in.min = simplify(apply_param_estimates(comp.in.min)); + comp.in.max = simplify(apply_param_estimates(comp.in.max)); if (equal(comp.in.min, req.min) && equal(comp.in.max, req.max)) { comp.equals_required = true; } else { @@ -911,11 +911,11 @@ FunctionDAG::FunctionDAG(const vector &outputs, const Target &target) stage.index = s; - exprs = apply_param_estimates.mutate(exprs); + exprs = apply_param_estimates(exprs); for (auto &p : func_value_bounds) { - p.second.min = apply_param_estimates.mutate(p.second.min); - p.second.max = apply_param_estimates.mutate(p.second.max); + p.second.min = apply_param_estimates(p.second.min); + p.second.max = apply_param_estimates(p.second.max); } // For this stage scope we want symbolic bounds for the rvars diff --git a/test/correctness/simd_op_check.h b/test/correctness/simd_op_check.h index 5a61df34252e..0fe23f1c1ef9 100644 --- a/test/correctness/simd_op_check.h +++ b/test/correctness/simd_op_check.h @@ -297,7 +297,7 @@ class SimdOpCheckTest { : image_params(image_params) { } } hook_up_image_params(image_params); - e = hook_up_image_params.mutate(e); + e = hook_up_image_params(e); class HasInlineReduction : public Internal::IRVisitor { using Internal::IRVisitor::visit; diff --git a/test/correctness/simd_op_check_sve2.cpp b/test/correctness/simd_op_check_sve2.cpp index f0183412323a..673c7c46949c 100644 --- a/test/correctness/simd_op_check_sve2.cpp +++ b/test/correctness/simd_op_check_sve2.cpp @@ -1266,7 +1266,7 @@ class SimdOpCheckArmSve : public SimdOpCheckTest { : env(env) { } } copier(env); - e = copier.mutate(e); + e = copier(e); } // Create Task and register diff --git a/test/correctness/specialize.cpp b/test/correctness/specialize.cpp index 1a807003f72a..dff81748aebc 100644 --- a/test/correctness/specialize.cpp +++ b/test/correctness/specialize.cpp @@ -441,7 +441,7 @@ int main(int argc, char **argv) { if_then_else_count = 0; CountIfThenElse pass1; for (auto ff : out.compile_to_module(out.infer_arguments()).functions()) { - pass1.mutate(ff.body); + pass1(ff.body); } Buffer input(3, 3), output(3, 3); @@ -471,8 +471,8 @@ int main(int argc, char **argv) { if_then_else_count = 0; CountIfThenElse pass2; - for (auto ff : out.compile_to_module(out.infer_arguments()).functions()) { - pass2.mutate(ff.body); + for (const auto &ff : out.compile_to_module(out.infer_arguments()).functions()) { + pass2(ff.body); } Buffer input(3, 3), output(3, 3);