diff --git a/src/AsyncProducers.cpp b/src/AsyncProducers.cpp index 1751f1d88221..3dd42a8a54d5 100644 --- a/src/AsyncProducers.cpp +++ b/src/AsyncProducers.cpp @@ -637,20 +637,12 @@ class TightenProducerConsumerNodes : public IRMutator { } return body; - } else if (const Block *block = body.as()) { + } else if (body.as()) { if (is_producer) { // We don't push produce nodes into blocks return ProducerConsumer::make(name, is_producer, body); } - vector sub_stmts; - Stmt rest; - do { - sub_stmts.push_back(block->first); - rest = block->rest; - block = rest.as(); - } while (block); - sub_stmts.push_back(rest); - + vector sub_stmts = Block::to_vector(body); for (Stmt &s : sub_stmts) { if (uses_vars.check_stmt(s)) { s = make_producer_consumer(name, is_producer, s, scope, uses_vars); @@ -813,14 +805,12 @@ class ExpandAcquireNodes : public IRMutator { Stmt visit(const Block *op) override { // Do an entire sequence of blocks in a single visit method to conserve stack space. - vector stmts; - Stmt result; - do { - stmts.push_back(mutate(op->first)); - result = op->rest; - } while ((op = result.as())); - - result = mutate(result); + vector stmts = Block::to_vector(op); + for (Stmt &s : stmts) { + s = mutate(s); + } + Stmt result = stmts.back(); + stmts.pop_back(); vector> semaphores; for (Stmt s : reverse_view(stmts)) { diff --git a/src/IR.cpp b/src/IR.cpp index 0af80a719969..8825c13f4c26 100644 --- a/src/IR.cpp +++ b/src/IR.cpp @@ -719,6 +719,20 @@ Stmt Block::make(const std::vector &stmts) { return result; } +std::vector Block::to_vector(const Stmt &s) { + std::vector result; + // Blocks are right-leaning: 'first' is never itself a Block. + Stmt rest = s; + while (const Block *b = rest.as()) { + result.push_back(b->first); + rest = b->rest; + } + if (rest.defined()) { + result.push_back(std::move(rest)); + } + return result; +} + Stmt Block::with(const Stmt &first, const Stmt &rest) const { if (first.same_as(this->first) && rest.same_as(this->rest)) { return this; diff --git a/src/IR.h b/src/IR.h index 6ee4515b717d..1c5f614d9c6c 100644 --- a/src/IR.h +++ b/src/IR.h @@ -602,6 +602,11 @@ struct Block : public StmtNode { * This method may not return a Block statement if stmts.size() <= 1. */ static Stmt make(const std::vector &stmts); + /** The inverse of the vector form of make. Unpacks a Stmt into the + * sequence of non-Block Stmts it runs. An undefined Stmt unpacks to an + * empty vector, and a non-Block Stmt unpacks to a vector of size one. */ + static std::vector to_vector(const Stmt &s); + static const IRNodeType _node_type = IRNodeType::Block; }; diff --git a/src/LoopCarry.cpp b/src/LoopCarry.cpp index fff1489ddf43..2817082ca034 100644 --- a/src/LoopCarry.cpp +++ b/src/LoopCarry.cpp @@ -104,24 +104,6 @@ class FindLoads : public IRGraphVisitor { vector result; }; -/** A helper for block_to_vector below. */ -void block_to_vector(const Stmt &s, vector &v) { - const Block *b = s.as(); - if (!b) { - v.push_back(s); - } else { - block_to_vector(b->first, v); - block_to_vector(b->rest, v); - } -} - -/** Unpack a block into its component Stmts. */ -vector block_to_vector(const Stmt &s) { - vector result; - block_to_vector(s, result); - return result; -} - Expr scratch_index(int i, Type t) { if (t.is_scalar()) { return i; @@ -221,7 +203,7 @@ class LoopCarryOverLoop : public IRMutator { } Stmt visit(const Block *op) override { - vector v = block_to_vector(op); + vector v = Block::to_vector(op); vector stores; vector result; diff --git a/src/Prefetch.cpp b/src/Prefetch.cpp index 71f9050902ed..51a204b6186e 100644 --- a/src/Prefetch.cpp +++ b/src/Prefetch.cpp @@ -390,43 +390,23 @@ class SplitPrefetch : public IRMutator { } }; -template -void traverse_block(const Stmt &s, Fn &&f) { - const Block *b = s.as(); - if (!b) { - f(s); - } else { - traverse_block(b->first, f); - traverse_block(b->rest, f); - } -} - class HoistPrefetches : public IRMutator { protected: using IRMutator::visit; Stmt visit(const Block *op) override { - Stmt s = op; - - Stmt prefetches, body; - traverse_block(s, [this, &prefetches, &body](const Stmt &s_in) { + vector prefetches, body; + for (const Stmt &s_in : Block::to_vector(op)) { Stmt s = IRMutator::mutate(s_in); const Evaluate *eval = s.as(); if (eval && Call::as_intrinsic(eval->value, {Call::prefetch})) { - prefetches = prefetches.defined() ? Block::make(prefetches, s) : s; + prefetches.push_back(std::move(s)); } else { - body = body.defined() ? Block::make(body, s) : s; + body.push_back(std::move(s)); } - }); - if (prefetches.defined()) { - if (body.defined()) { - return Block::make(prefetches, body); - } else { - return prefetches; - } - } else { - return body; } + prefetches.insert(prefetches.end(), body.begin(), body.end()); + return Block::make(prefetches); } }; diff --git a/src/RemoveUndef.cpp b/src/RemoveUndef.cpp index 9a1f0635d6cc..b421cdfc86d7 100644 --- a/src/RemoveUndef.cpp +++ b/src/RemoveUndef.cpp @@ -531,30 +531,24 @@ class RemoveUndef : public IRMutator { Stmt visit(const Block *op) override { // Visit a sequence of blocks in a single method to conserve stack space. - Stmt result; - vector> frames; - - do { - Stmt next = mutate(op->first); - if (next.defined()) { - frames.emplace_back(op, std::move(next)); + vector stmts = Block::to_vector(op); + vector new_stmts; + new_stmts.reserve(stmts.size()); + bool unchanged = true; + for (const Stmt &s : stmts) { + Stmt new_s = mutate(s); + unchanged &= new_s.same_as(s); + if (new_s.defined()) { + new_stmts.push_back(std::move(new_s)); } - result = op->rest; - } while ((op = result.as())); - - result = mutate(result); + } - for (const auto &[block, stmt] : reverse_view(frames)) { - Stmt new_first = stmt; - if (!result.defined()) { - result = new_first; - } else if (new_first.same_as(block->first) && result.same_as(block->rest)) { - result = block; - } else { - result = Block::make(new_first, result); - } + if (unchanged) { + return op; + } else { + // Returns an undefined Stmt if everything was removed. + return Block::make(new_stmts); } - return result; } Stmt visit(const IfThenElse *op) override { diff --git a/src/StmtToHTML.cpp b/src/StmtToHTML.cpp index 030c56618e92..b0fae24e6eae 100644 --- a/src/StmtToHTML.cpp +++ b/src/StmtToHTML.cpp @@ -1262,11 +1262,8 @@ class HTMLCodePrinter : public IRVisitor { // To avoid generating ridiculously deep DOMs, we flatten blocks here. void print_block_stmt(const Stmt &stmt) { - if (const Block *b = stmt.as()) { - print_block_stmt(b->first); - print_block_stmt(b->rest); - } else if (stmt.defined()) { - print(stmt); + for (const Stmt &s : Block::to_vector(stmt)) { + print(s); } }