@samitouri / QOS-React-2 / commits / de0c2393ad

[rust] MergeConsecutiveBlocks and block helpers

Ports `MergeConsecutiveBlocks` to Rust. This was a tricky one: as we iterate through the blocks _if_ the block ends up being merged with its predecessor we need to consume it and modify its predecessor block (ie, mutating two things from the same data structure - shared mutation!). But if we _don't_ need to merge, then we need to not drop the current block. Ie, we sort of need to conditionally take ownership of the current block during iteration and put it back. I added a `BlockRewriter` helper type for this which has a helper to iterate safely. It calls the iterator lambda, moving blocks one at a time into the lambda. The lambda returns either `Keep(block)` to give the block back and keep it or `Remove` to tell the rewriter to drop the block. Thanks to making the Blocks data structure hold `Option<Box<BasicBlock>>` items, "moving" the block actually just means nulling out the option and conceptually giving ownership of the pointer to the callback - the data itself never moves. This involved creating a custom `Blocks` wrapper type, which i had been putting off doing. This cleans up a bunch of other logic around traversing blocks.

Joe Savona committed Jul 14, 2023 at 18:00 UTC de0c2393add8ae48b107774fb87d8b4a04c630ed
14 files changed +410 -143
compiler/forget/Cargo.lock
+1
@@ -535,6 +535,7 @@ version = "0.1.0"
535 dependencies = [
536 "bumpalo",
537 "forget_build_hir",
538 + "forget_diagnostics",
539 "forget_estree",
540 "forget_hir",
541 "forget_ssa",
compiler/forget/crates/forget_build_hir/src/builder.rs
+24 -32
@@ -1,15 +1,13 @@
1 use std::cell::RefCell;
2 use std::rc::Rc;
3
4 -use bumpalo::boxed::Box;
4 use bumpalo::collections::{String, Vec};
5 use forget_diagnostics::Diagnostic;
6 use forget_hir::{
8 - initialize_hir, BasicBlock, BlockId, BlockKind, Environment, GotoKind, Identifier,
7 + initialize_hir, BasicBlock, BlockId, BlockKind, Blocks, Environment, GotoKind, Identifier,
8 IdentifierData, InstrIx, Instruction, InstructionIdGenerator, InstructionValue, Terminal,
9 TerminalValue, Type, HIR,
10 };
12 -use indexmap::IndexMap;
11
12 use crate::BuildHIRError;
13
@@ -26,7 +24,7 @@ pub(crate) struct Builder<'a> {
24 #[allow(dead_code)]
25 environment: &'a Environment<'a>,
26
29 - completed: IndexMap<BlockId, Box<'a, BasicBlock<'a>>>,
27 + completed: Blocks<'a>,
28
29 instructions: Vec<'a, Instruction<'a>>,
30
@@ -153,20 +151,17 @@ impl<'a> Builder<'a> {
151 fallthrough: WipBlock<'a>,
152 ) {
153 let prev_wip = std::mem::replace(&mut self.wip, fallthrough);
156 - self.completed.insert(
157 - prev_wip.id,
158 - self.environment.box_new(BasicBlock {
159 - id: prev_wip.id,
160 - kind: prev_wip.kind,
161 - instructions: prev_wip.instructions,
162 - terminal: Terminal {
163 - id: self.id_gen.next(),
164 - value: terminal,
165 - },
166 - predecessors: Default::default(),
167 - phis: self.environment.vec_new(),
168 - }),
169 - );
154 + self.completed.insert(Box::new(BasicBlock {
155 + id: prev_wip.id,
156 + kind: prev_wip.kind,
157 + instructions: prev_wip.instructions,
158 + terminal: Terminal {
159 + id: self.id_gen.next(),
160 + value: terminal,
161 + },
162 + predecessors: Default::default(),
163 + phis: self.environment.vec_new(),
164 + }));
165 }
166
167 pub(crate) fn reserve(&mut self, kind: BlockKind) -> WipBlock<'a> {
@@ -206,20 +201,17 @@ impl<'a> Builder<'a> {
201 };
202
203 let completed = std::mem::replace(&mut self.wip, current);
209 - self.completed.insert(
210 - completed.id,
211 - self.environment.box_new(BasicBlock {
212 - id: completed.id,
213 - kind: completed.kind,
214 - instructions: completed.instructions,
215 - terminal: Terminal {
216 - id: self.id_gen.next(),
217 - value: terminal,
218 - },
219 - predecessors: Default::default(),
220 - phis: self.environment.vec_new(),
221 - }),
222 - );
204 + self.completed.insert(Box::new(BasicBlock {
205 + id: completed.id,
206 + kind: completed.kind,
207 + instructions: completed.instructions,
208 + terminal: Terminal {
209 + id: self.id_gen.next(),
210 + value: terminal,
211 + },
212 + predecessors: Default::default(),
213 + phis: self.environment.vec_new(),
214 + }));
215 result
216 }
217
compiler/forget/crates/forget_fixtures/tests/fixtures_test.rs
+1 -1
@@ -46,7 +46,7 @@ fn fixtures() {
46 println!("ok enter_ssa");
47 eliminate_redundant_phis(&environment, &mut fun);
48 println!("ok eliminate_redundant_phis");
49 - constant_propagation(&environment, &mut fun);
49 + constant_propagation(&environment, &mut fun).unwrap();
50 println!("ok constant_propagation");
51 fun.print(&fun.body, &mut output).unwrap();
52 println!("ok print");
compiler/forget/crates/forget_fixtures/tests/snapshots/fixtures_test__fixtures@constant-propagation-constant-if-condition.js.snap
+20 -44
@@ -51,47 +51,23 @@ bb0 (block)
51 [3] #7 = 1
52 [4] #8 = 1
53 [5] #9 = true
54 - [6] Goto bb2
55 -bb2 (block)
56 - predecessors: bb0
57 - [7] #3 = true
58 - [8] #4 = StoreLocal Reassign unknown b$7 = unknown #3
59 - [9] Goto bb1
60 -bb1 (block)
61 - predecessors: bb2
62 - [10] #10 = DeclareLocal Let unknown c$9
63 - [11] #15 = true
64 - [12] Goto bb5
65 -bb5 (block)
66 - predecessors: bb1
67 - [13] #11 = "hello"
68 - [14] #12 = StoreLocal Reassign unknown c$11 = unknown #11
69 - [15] Goto bb4
70 -bb4 (block)
71 - predecessors: bb5
72 - [16] #16 = DeclareLocal Let unknown d$13
73 - [17] #21 = "hello"
74 - [18] #22 = "hello"
75 - [19] #23 = true
76 - [20] Goto bb8
77 -bb8 (block)
78 - predecessors: bb4
79 - [21] #17 = 42
80 - [22] #18 = StoreLocal Reassign unknown d$15 = unknown #17
81 - [23] Goto bb7
82 -bb7 (block)
83 - predecessors: bb8
84 - [24] #24 = DeclareLocal Let unknown e$17
85 - [25] #29 = 42
86 - [26] #30 = 42
87 - [27] #31 = true
88 - [28] Goto bb11
89 -bb11 (block)
90 - predecessors: bb7
91 - [29] #25 = "ok"
92 - [30] #26 = StoreLocal Reassign unknown e$19 = unknown #25
93 - [31] Goto bb10
94 -bb10 (block)
95 - predecessors: bb11
96 - [32] #32 = "ok"
97 - [33] Return unknown #32
54 + [6] #3 = true
55 + [7] #4 = StoreLocal Reassign unknown b$7 = unknown #3
56 + [8] #10 = DeclareLocal Let unknown c$9
57 + [9] #15 = true
58 + [10] #11 = "hello"
59 + [11] #12 = StoreLocal Reassign unknown c$11 = unknown #11
60 + [12] #16 = DeclareLocal Let unknown d$13
61 + [13] #21 = "hello"
62 + [14] #22 = "hello"
63 + [15] #23 = true
64 + [16] #17 = 42
65 + [17] #18 = StoreLocal Reassign unknown d$15 = unknown #17
66 + [18] #24 = DeclareLocal Let unknown e$17
67 + [19] #29 = 42
68 + [20] #30 = 42
69 + [21] #31 = true
70 + [22] #25 = "ok"
71 + [23] #26 = StoreLocal Reassign unknown e$19 = unknown #25
72 + [24] #32 = "ok"
73 + [25] Return unknown #32
compiler/forget/crates/forget_fixtures/tests/snapshots/fixtures_test__fixtures@function-expressions.js.snap
+17 -23
@@ -47,24 +47,18 @@ bb0 (block)
47 [3] #9 = 1
48 [4] #10 = 1
49 [5] #11 = true
50 - [6] Goto bb3
51 - bb3 (block)
52 - predecessors: bb1
53 - [7] #3 = 5
54 - [8] #4 = 3
55 - [9] #5 = 8
56 - [10] #6 = StoreLocal Reassign unknown b$16 = unknown #5
57 - [11] Goto bb2
58 - bb2 (block)
59 - predecessors: bb3
60 - [12] #12 = 2
61 - [13] #13 = LoadLocal unknown y$13
62 - [14] #14 = Binary unknown #12 + unknown #13
63 - [15] #15 = 1
64 - [16] #16 = Binary unknown #14 + unknown #15
65 - [17] #17 = 8
66 - [18] #18 = Binary unknown #16 + unknown #17
67 - [19] #19 = Function @deps[] @context[unknown x$18, unknown y$19, unknown a$20, unknown b$21]:
50 + [6] #3 = 5
51 + [7] #4 = 3
52 + [8] #5 = 8
53 + [9] #6 = StoreLocal Reassign unknown b$16 = unknown #5
54 + [10] #12 = 2
55 + [11] #13 = LoadLocal unknown y$13
56 + [12] #14 = Binary unknown #12 + unknown #13
57 + [13] #15 = 1
58 + [14] #16 = Binary unknown #14 + unknown #15
59 + [15] #17 = 8
60 + [16] #18 = Binary unknown #16 + unknown #17
61 + [17] #19 = Function @deps[] @context[unknown x$18, unknown y$19, unknown a$20, unknown b$21]:
62 function bar(
63 unknown z$22,
64 )
@@ -90,11 +84,11 @@ bb0 (block)
84 [17] #17 = Binary unknown #15 + unknown #16
85 [18] #18 = <undefined>
86 [19] Return unknown #18
93 - [20] #20 = StoreLocal Const unknown bar$26 = unknown #19
94 - [21] #21 = LoadLocal unknown bar$26
95 - [22] #22 = LoadLocal unknown foo$1
96 - [23] #23 = <undefined>
97 - [24] Return unknown #23
87 + [18] #20 = StoreLocal Const unknown bar$26 = unknown #19
88 + [19] #21 = LoadLocal unknown bar$26
89 + [20] #22 = LoadLocal unknown foo$1
90 + [21] #23 = <undefined>
91 + [22] Return unknown #23
92 [3] #3 = StoreLocal Const unknown foo$28 = unknown #2
93 [4] #4 = <undefined>
94 [5] Return unknown #4
compiler/forget/crates/forget_hir/src/function.rs
+143 -5
@@ -1,5 +1,5 @@
1 -use bumpalo::boxed::Box;
1 use bumpalo::collections::{String, Vec};
2 +use forget_diagnostics::Diagnostic;
3 use indexmap::IndexMap;
4
5 use crate::{BasicBlock, BlockId, IdentifierOperand, Instruction};
@@ -32,14 +32,152 @@ pub struct HIR<'a> {
32 pub instructions: Vec<'a, Instruction<'a>>,
33 }
34
35 -pub type Blocks<'a> = IndexMap<BlockId, Box<'a, BasicBlock<'a>>>;
35 +#[derive(Default, Debug)]
36 +pub struct Blocks<'a> {
37 + data: IndexMap<BlockId, Option<Box<BasicBlock<'a>>>>,
38 +}
39 +
40 +impl<'a> Blocks<'a> {
41 + pub fn new() -> Self {
42 + Self {
43 + data: Default::default(),
44 + }
45 + }
46 +
47 + pub fn with_capacity(capacity: usize) -> Self {
48 + Self {
49 + data: IndexMap::with_capacity(capacity),
50 + }
51 + }
52 +
53 + pub fn len(&self) -> usize {
54 + self.data.len()
55 + }
56 +
57 + pub fn insert(&mut self, block: Box<BasicBlock<'a>>) -> Option<Option<Box<BasicBlock<'a>>>> {
58 + self.data.insert(block.id, Some(block))
59 + }
60 +
61 + pub fn block_ids(&self) -> std::vec::Vec<BlockId> {
62 + self.data.keys().cloned().collect()
63 + }
64 +
65 + pub fn take(&mut self, id: BlockId) -> Box<BasicBlock<'a>> {
66 + self.data.remove(&id).unwrap().unwrap()
67 + }
68
37 -impl<'a> HIR<'a> {
69 pub fn block(&self, id: BlockId) -> &BasicBlock<'a> {
39 - self.blocks.get(&id).unwrap()
70 + self.data.get(&id).unwrap().as_ref().unwrap()
71 }
72
73 pub fn block_mut(&mut self, id: BlockId) -> &mut BasicBlock<'a> {
43 - self.blocks.get_mut(&id).unwrap()
74 + self.data.get_mut(&id).unwrap().as_mut().unwrap()
75 + }
76 +
77 + pub fn iter(&self) -> BlocksIter<'_, 'a> {
78 + BlocksIter::new(self.data.iter())
79 + }
80 +
81 + pub fn iter_mut(&mut self) -> BlocksIterMut<'_, 'a> {
82 + BlocksIterMut::new(self.data.iter_mut())
83 + }
84 +}
85 +
86 +pub struct BlocksIter<'b, 'a> {
87 + iter: indexmap::map::Iter<'b, BlockId, Option<Box<BasicBlock<'a>>>>,
88 +}
89 +
90 +impl<'b, 'a> BlocksIter<'b, 'a> {
91 + fn new(iter: indexmap::map::Iter<'b, BlockId, Option<Box<BasicBlock<'a>>>>) -> Self {
92 + Self { iter }
93 + }
94 +}
95 +
96 +impl<'b, 'a> Iterator for BlocksIter<'b, 'a> {
97 + type Item = &'b BasicBlock<'a>;
98 +
99 + fn size_hint(&self) -> (usize, Option<usize>) {
100 + self.iter.size_hint()
101 }
102 +
103 + fn next(&mut self) -> Option<Self::Item> {
104 + self.iter
105 + .next()
106 + .and_then(|(_, block)| block.as_ref())
107 + .map(|block| block.as_ref())
108 + }
109 +}
110 +
111 +pub struct BlocksIterMut<'b, 'a> {
112 + iter: indexmap::map::IterMut<'b, BlockId, Option<Box<BasicBlock<'a>>>>,
113 +}
114 +
115 +impl<'b, 'a> BlocksIterMut<'b, 'a> {
116 + fn new(iter: indexmap::map::IterMut<'b, BlockId, Option<Box<BasicBlock<'a>>>>) -> Self {
117 + Self { iter }
118 + }
119 +}
120 +
121 +impl<'b, 'a> Iterator for BlocksIterMut<'b, 'a> {
122 + type Item = &'b mut BasicBlock<'a>;
123 +
124 + fn size_hint(&self) -> (usize, Option<usize>) {
125 + self.iter.size_hint()
126 + }
127 +
128 + fn next(&mut self) -> Option<Self::Item> {
129 + self.iter
130 + .next()
131 + .and_then(|(_, block)| block.as_mut())
132 + .map(|block| block.as_mut())
133 + }
134 +}
135 +
136 +pub struct BlockRewriter<'blocks, 'a> {
137 + blocks: &'blocks mut Blocks<'a>,
138 + current: BlockId,
139 +}
140 +
141 +impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
142 + pub fn new(blocks: &'blocks mut Blocks<'a>, entry: BlockId) -> Self {
143 + Self {
144 + blocks,
145 + current: entry,
146 + }
147 + }
148 +
149 + pub fn each_block<F>(&mut self, mut f: F) -> Result<(), Diagnostic>
150 + where
151 + F: FnMut(Box<BasicBlock<'a>>, &mut Self) -> Result<BlockRewriterAction<'a>, Diagnostic>,
152 + {
153 + let keys = self.blocks.block_ids();
154 + for block_id in keys {
155 + self.current = block_id;
156 + let block = self.blocks.data.get_mut(&block_id).unwrap().take().unwrap();
157 + match f(block, self)? {
158 + BlockRewriterAction::Keep(block) => {
159 + self.blocks.data.insert(block_id, Some(block));
160 + }
161 + BlockRewriterAction::Remove => {
162 + // nothing to do, already removed from the blocks
163 + }
164 + }
165 + }
166 + Ok(())
167 + }
168 +
169 + pub fn block(&self, block_id: BlockId) -> &BasicBlock<'a> {
170 + assert_ne!(block_id, self.current);
171 + self.blocks.block(block_id)
172 + }
173 +
174 + pub fn block_mut(&mut self, block_id: BlockId) -> &mut BasicBlock<'a> {
175 + assert_ne!(block_id, self.current);
176 + self.blocks.block_mut(block_id)
177 + }
178 +}
179 +
180 +pub enum BlockRewriterAction<'a> {
181 + Keep(Box<BasicBlock<'a>>),
182 + Remove,
183 }
compiler/forget/crates/forget_hir/src/initialize.rs
+13 -14
@@ -1,10 +1,9 @@
1 use std::collections::HashSet;
2
3 use forget_diagnostics::{invariant, Diagnostic};
4 -use indexmap::IndexMap;
4 use thiserror::Error;
5
7 -use crate::{BlockId, GotoKind, GotoTerminal, InstructionIdGenerator, TerminalValue, HIR};
6 +use crate::{BlockId, Blocks, GotoKind, GotoTerminal, InstructionIdGenerator, TerminalValue, HIR};
7
8 /// Runs a variety of passes to put the HIR in canonical form. This should be called
9 /// after initial HIR construction and after any transformations that change the
@@ -34,7 +33,7 @@ pub fn reverse_postorder_blocks<'a>(hir: &mut HIR<'a>) {
33 // already visited
34 return;
35 }
37 - let block = hir.block(block_id);
36 + let block = hir.blocks.block(block_id);
37 let terminal = &block.terminal;
38 match &terminal.value {
39 TerminalValue::Branch(terminal) => {
@@ -64,9 +63,9 @@ pub fn reverse_postorder_blocks<'a>(hir: &mut HIR<'a>) {
63 visit(hir.entry, &hir, &mut visited, &mut postorder);
64
65 // NOTE: could consider sorting the blocks in-place by key
67 - let mut blocks = IndexMap::with_capacity(hir.blocks.len());
66 + let mut blocks = Blocks::with_capacity(hir.blocks.len());
67 for id in postorder.iter().rev().cloned() {
69 - blocks.insert(id, hir.blocks.remove(&id).unwrap());
68 + blocks.insert(hir.blocks.take(id));
69 }
70
71 hir.blocks = blocks;
@@ -74,9 +73,9 @@ pub fn reverse_postorder_blocks<'a>(hir: &mut HIR<'a>) {
73
74 /// Prunes ForTerminal.update values (sets to None) if they are unreachable
75 pub fn remove_unreachable_for_updates<'a>(hir: &mut HIR<'a>) {
77 - let block_ids: HashSet<BlockId> = hir.blocks.keys().cloned().collect();
76 + let block_ids = hir.blocks.block_ids();
77
79 - for block in hir.blocks.values_mut() {
78 + for block in hir.blocks.iter_mut() {
79 if let TerminalValue::For(terminal) = &mut block.terminal.value {
80 if let Some(update) = terminal.update {
81 if !block_ids.contains(&update) {
@@ -90,9 +89,9 @@ pub fn remove_unreachable_for_updates<'a>(hir: &mut HIR<'a>) {
89 /// Prunes unreachable fallthrough values, setting them to None if the referenced
90 /// block was not otherwise reachable.
91 pub fn remove_unreachable_fallthroughs<'a>(hir: &mut HIR<'a>) {
93 - let block_ids: HashSet<BlockId> = hir.blocks.keys().cloned().collect();
92 + let block_ids = hir.blocks.block_ids();
93
95 - for block in hir.blocks.values_mut() {
94 + for block in hir.blocks.iter_mut() {
95 block
96 .terminal
97 .value
@@ -108,9 +107,9 @@ pub fn remove_unreachable_fallthroughs<'a>(hir: &mut HIR<'a>) {
107
108 /// Rewrites DoWhile statements into Gotos if the test block is not reachable
109 pub fn remove_unreachable_do_while_statements<'a>(hir: &mut HIR<'a>) {
111 - let block_ids: HashSet<BlockId> = hir.blocks.keys().cloned().collect();
110 + let block_ids = hir.blocks.block_ids();
111
113 - for block in hir.blocks.values_mut() {
112 + for block in hir.blocks.iter_mut() {
113 if let TerminalValue::DoWhile(terminal) = &mut block.terminal.value {
114 if !block_ids.contains(&terminal.test) {
115 block.terminal.value = TerminalValue::Goto(GotoTerminal {
@@ -127,7 +126,7 @@ pub fn remove_unreachable_do_while_statements<'a>(hir: &mut HIR<'a>) {
126 pub fn mark_instruction_ids<'a>(hir: &mut HIR<'a>) -> Result<(), Diagnostic> {
127 let mut id_gen = InstructionIdGenerator::new();
128 let mut visited = HashSet::<(usize, usize)>::new();
130 - for (ii, block) in hir.blocks.values_mut().enumerate() {
129 + for (ii, block) in hir.blocks.iter_mut().enumerate() {
130 let block_id = block.id;
131 for (jj, instr_ix) in block.instructions.iter_mut().enumerate() {
132 invariant(visited.insert((ii, jj)), || {
@@ -149,7 +148,7 @@ pub struct BlockVisitedTwice {
148
149 /// Updates the predecessors of each block
150 pub fn mark_predecessors<'a>(hir: &mut HIR<'a>) {
152 - for block in hir.blocks.values_mut() {
151 + for block in hir.blocks.iter_mut() {
152 block.predecessors.clear();
153 }
154 let mut visited = HashSet::<BlockId>::with_capacity(hir.blocks.len());
@@ -159,7 +158,7 @@ pub fn mark_predecessors<'a>(hir: &mut HIR<'a>) {
158 hir: &mut HIR<'a>,
159 visited: &mut HashSet<BlockId>,
160 ) {
162 - let block = hir.block_mut(block_id);
161 + let block = hir.blocks.block_mut(block_id);
162 if let Some(prev_id) = prev_id {
163 block.predecessors.insert(prev_id);
164 }
compiler/forget/crates/forget_hir/src/lib.rs
+2
@@ -5,6 +5,7 @@ mod function;
5 mod id_types;
6 mod initialize;
7 mod instruction;
8 +mod merge_consecutive_blocks;
9 mod print;
10 mod registry;
11 mod terminal;
@@ -21,6 +22,7 @@ pub use initialize::{
22 remove_unreachable_for_updates, reverse_postorder_blocks,
23 };
24 pub use instruction::*;
25 +pub use merge_consecutive_blocks::merge_consecutive_blocks;
26 pub use print::Print;
27 pub use registry::Registry;
28 pub use terminal::*;
compiler/forget/crates/forget_hir/src/merge_consecutive_blocks.rs new
+152
@@ -0,0 +1,152 @@
1 +use std::collections::HashMap;
2 +
3 +use forget_diagnostics::{invariant, Diagnostic};
4 +use thiserror::Error;
5 +
6 +use crate::{
7 + mark_instruction_ids, mark_predecessors, BasicBlock, BlockId, BlockKind, BlockRewriter,
8 + BlockRewriterAction, Environment, Function, IdentifierOperand, InstrIx, Instruction,
9 + InstructionKind, InstructionValue, LValue, LoadLocal, Operand, StoreLocal, TerminalValue,
10 +};
11 +
12 +/// Merges sequences of blocks that will always execute consecutively —
13 +/// ie where the predecessor always transfers control to the successor
14 +/// (ends in a goto) and where the predecessor is the only predecessor
15 +/// for that successor (ie, there is no other way to reach the successor).
16 +///
17 +/// Note that this pass leaves value/loop blocks alone because they cannot
18 +/// be merged without breaking the structure of the high-level terminals
19 +/// that reference them.
20 +pub fn merge_consecutive_blocks<'a>(
21 + env: &Environment<'a>,
22 + fun: &mut Function<'a>,
23 +) -> Result<(), Diagnostic> {
24 + let mut merged = MergedBlocks::default();
25 + let blocks = &mut fun.body.blocks;
26 + let instructions = &mut fun.body.instructions;
27 + let mut rewriter = BlockRewriter::new(blocks, fun.body.entry);
28 + let mut has_changes = false;
29 +
30 + rewriter.each_block(|mut block, rewriter| {
31 + let block_id = block.id;
32 + // Visit instructions to merge blocks within function expressions
33 + for instr_ix in &block.instructions {
34 + let instr = &mut instructions[usize::from(*instr_ix)];
35 + if let InstructionValue::Function(fun) = &mut instr.value {
36 + merge_consecutive_blocks(env, &mut fun.lowered_function)?;
37 + }
38 + }
39 +
40 + // Can't merge value blocks and can't merge blocks with multiple
41 + // predecessors
42 + if block.kind != BlockKind::Block || block.predecessors.len() != 1 {
43 + return Ok(BlockRewriterAction::Keep(block));
44 + }
45 +
46 + let original_predecessor_id = block.predecessors.first().unwrap(); // length checked above
47 + let predecessor_id = merged.get(*original_predecessor_id);
48 + let predecessor = rewriter.block_mut(predecessor_id);
49 + if predecessor.kind != BlockKind::Block
50 + || !matches!(predecessor.terminal.value, TerminalValue::Goto(_))
51 + {
52 + // Can't merge value blocks, and we can't merge if the predecessor
53 + // has multiple successors (and isn't guaranteed to transfer here)
54 + return Ok(BlockRewriterAction::Keep(block));
55 + }
56 +
57 + // Replace phis in the merged block with canonical assignments to the single
58 + // operand value
59 + for phi in block.phis.iter_mut() {
60 + invariant(phi.operands.len() == 1, || {
61 + Diagnostic::invariant(ExpectedSingleOperandPhis { block: block_id }, None)
62 + })?;
63 + let (_, operand) = phi.operands.first().unwrap();
64 + // load the operand
65 + let load = Instruction {
66 + id: predecessor.terminal.id,
67 + value: InstructionValue::LoadLocal(LoadLocal {
68 + place: IdentifierOperand {
69 + effect: None,
70 + identifier: operand.clone(),
71 + },
72 + }),
73 + };
74 + let load_ix = InstrIx::new(instructions.len() as u32);
75 + instructions.push(load);
76 + predecessor.instructions.push(load_ix);
77 + // store it into the phi id
78 + let store = Instruction {
79 + id: predecessor.terminal.id,
80 + value: InstructionValue::StoreLocal(StoreLocal {
81 + lvalue: LValue {
82 + kind: InstructionKind::Reassign,
83 + identifier: IdentifierOperand {
84 + identifier: phi.identifier.clone(),
85 + effect: None,
86 + },
87 + },
88 + value: Operand {
89 + effect: None,
90 + ix: load_ix,
91 + },
92 + }),
93 + };
94 + let store_ix = InstrIx::new(instructions.len() as u32);
95 + instructions.push(store);
96 + predecessor.instructions.push(store_ix);
97 + }
98 + let BasicBlock {
99 + instructions,
100 + terminal,
101 + ..
102 + } = *block;
103 + predecessor.instructions.extend(instructions);
104 + predecessor.terminal = terminal;
105 + merged.merge(block_id, predecessor_id);
106 +
107 + has_changes = true;
108 + Ok(BlockRewriterAction::Remove)
109 + })?;
110 +
111 + if has_changes {
112 + mark_instruction_ids(&mut fun.body)?;
113 + mark_predecessors(&mut fun.body);
114 + }
115 +
116 + Ok(())
117 +}
118 +
119 +#[derive(Default)]
120 +struct MergedBlocks {
121 + merged: HashMap<BlockId, BlockId>,
122 +}
123 +
124 +impl MergedBlocks {
125 + fn merge(&mut self, block: BlockId, into: BlockId) {
126 + let target = self.get(into);
127 + self.merged.insert(block, target);
128 + }
129 +
130 + fn get(&self, block: BlockId) -> BlockId {
131 + let mut current = block;
132 + while let Some(mapped) = self.merged.get(&current) {
133 + current = *mapped;
134 + }
135 + current
136 + }
137 +}
138 +
139 +#[derive(Debug, Error)]
140 +#[error("Expected predecessor {predecessor} to exist")]
141 +pub struct ExpectedPredecessorToExist {
142 + predecessor: BlockId,
143 +}
144 +
145 +#[derive(Debug, Error)]
146 +#[error(
147 + "Expected block {block} with single predecessor to have no phis or
148 + phis with a single operand, found multiple operands"
149 +)]
150 +pub struct ExpectedSingleOperandPhis {
151 + block: BlockId,
152 +}
compiler/forget/crates/forget_hir/src/print.rs
+1 -1
@@ -34,7 +34,7 @@ impl<'a> Print<'a> for Function<'a> {
34 }
35 writeln!(out, ")")?;
36 writeln!(out, "entry {}", self.body.entry)?;
37 - for (_, block) in self.body.blocks.iter() {
37 + for block in self.body.blocks.iter() {
38 block.print(hir, out)?;
39 }
40 writeln!(out)?;
compiler/forget/crates/forget_optimization/Cargo.toml
+1
@@ -18,6 +18,7 @@ forget_hir = { workspace = true }
18 forget_ssa = { workspace = true }
19 forget_build_hir = { workspace = true }
20 forget_utils = { workspace = true }
21 +forget_diagnostics = { workspace = true }
22 bumpalo = { workspace = true }
23 indexmap = { workspace = true }
24 miette = { workspace = true }
compiler/forget/crates/forget_optimization/src/constant_propagation.rs
+24 -13
@@ -1,24 +1,29 @@
1 use std::collections::HashMap;
2
3 +use forget_diagnostics::Diagnostic;
4 use forget_estree::BinaryOperator;
5 use forget_hir::{
5 - initialize_hir, BlockKind, Environment, Function, GotoKind, GotoTerminal, IdentifierId,
6 - Instruction, InstructionValue, LoadGlobal, Operand, Primitive, PrimitiveValue, TerminalValue,
6 + initialize_hir, merge_consecutive_blocks, BlockKind, Environment, Function, GotoKind,
7 + GotoTerminal, IdentifierId, Instruction, InstructionValue, LoadGlobal, Operand, Primitive,
8 + PrimitiveValue, TerminalValue,
9 };
10 use forget_ssa::eliminate_redundant_phis;
11
10 -pub fn constant_propagation<'a>(env: &Environment<'a>, fun: &mut Function<'a>) {
12 +pub fn constant_propagation<'a>(
13 + env: &Environment<'a>,
14 + fun: &mut Function<'a>,
15 +) -> Result<(), Diagnostic> {
16 let mut constants = Constants::new();
12 - constant_propagation_impl(env, fun, &mut constants);
17 + constant_propagation_impl(env, fun, &mut constants)
18 }
19
20 fn constant_propagation_impl<'a>(
21 env: &Environment<'a>,
22 fun: &mut Function<'a>,
23 constants: &mut Constants<'a>,
19 -) {
24 +) -> Result<(), Diagnostic> {
25 loop {
21 - let have_terminals_changed = apply_constant_propagation(env, fun, constants);
26 + let have_terminals_changed = apply_constant_propagation(env, fun, constants)?;
27 if !have_terminals_changed {
28 break;
29 }
@@ -30,7 +35,7 @@ fn constant_propagation_impl<'a>(
35 // Now that predecessors have changed, prune phi operands for unreachable blocks
36 // for example, a phi node whose operand was eliminated because it was set in a
37 // block that is no longer reached
33 - for (_, block) in fun.body.blocks.iter_mut() {
38 + for block in fun.body.blocks.iter_mut() {
39 // TODO: avoid the clone here
40 let predecessors = block.predecessors.clone();
41 for phi in block.phis.iter_mut() {
@@ -42,17 +47,22 @@ fn constant_propagation_impl<'a>(
47 // By removing some phi operands, there may be phis that were not previously
48 // redundant but now are
49 eliminate_redundant_phis(env, fun);
50 +
51 + // Finally, merge together any blocks that are now guaranteed to execute
52 + // consecutively
53 + merge_consecutive_blocks(env, fun)?;
54 }
55 + Ok(())
56 }
57
58 fn apply_constant_propagation<'a>(
59 env: &Environment<'a>,
60 fun: &mut Function<'a>,
61 constants: &mut Constants<'a>,
52 -) -> bool {
62 +) -> Result<bool, Diagnostic> {
63 let mut has_changes = false;
64
55 - for (_, block) in fun.body.blocks.iter_mut() {
65 + for block in fun.body.blocks.iter_mut() {
66 for phi in block.phis.iter() {
67 let mut value: Option<Constant<'a>> = None;
68 for (_, operand) in &phi.operands {
@@ -90,7 +100,7 @@ fn apply_constant_propagation<'a>(
100 &mut fun.body.instructions[instr_ix].value,
101 InstructionValue::Tombstone,
102 );
93 - evaluate_instruction(env, &fun.body.instructions, &mut instr, constants);
103 + evaluate_instruction(env, &fun.body.instructions, &mut instr, constants)?;
104 fun.body.instructions[instr_ix].value = instr;
105 }
106
@@ -115,7 +125,7 @@ fn apply_constant_propagation<'a>(
125 }
126 }
127
118 - has_changes
128 + Ok(has_changes)
129 }
130
131 fn read_primitive_instruction<'a>(
@@ -135,7 +145,7 @@ fn evaluate_instruction<'a>(
145 instrs: &[Instruction<'a>],
146 mut instr: &mut InstructionValue<'a>,
147 constants: &mut Constants<'a>,
138 -) {
148 +) -> Result<(), Diagnostic> {
149 let read_constant = |operand: &Operand| {
150 let instr = &instrs[usize::from(operand.ix)].value;
151 match instr {
@@ -188,12 +198,13 @@ fn evaluate_instruction<'a>(
198 value.map(|value| (id.identifier.id, value.clone()))
199 })
200 .collect();
191 - constant_propagation_impl(env, &mut value.lowered_function, &mut inner_constants);
201 + constant_propagation_impl(env, &mut value.lowered_function, &mut inner_constants)?;
202 }
203 _ => {
204 // no-op, not all instructions can be processed
205 }
206 }
207 + Ok(())
208 }
209
210 fn apply_binary_operator<'a>(
compiler/forget/crates/forget_ssa/src/eliminate_redundant_phis.rs
+1 -1
@@ -28,7 +28,7 @@ pub fn eliminate_redundant_phis<'a>(env: &Environment, fun: &mut Function<'a>) {
28 let is_first_iteration = !has_back_edge;
29 len = rewrites.len();
30
31 - for (_, block) in hir.blocks.iter_mut() {
31 + for block in hir.blocks.iter_mut() {
32 if !has_back_edge {
33 for predecessor in block.predecessors.iter() {
34 if !visited.contains(predecessor) {
compiler/forget/crates/forget_ssa/src/enter.rs
+10 -9
@@ -1,7 +1,7 @@
1 use std::cell::RefCell;
2 use std::rc::Rc;
3
4 -use bumpalo::collections::{CollectIn, Vec};
4 +use bumpalo::collections::Vec;
5 use forget_diagnostics::{invariant, Diagnostic};
6 use forget_hir::{
7 BasicBlock, BlockId, Blocks, Environment, Function, Identifier, IdentifierData, IdentifierId,
@@ -32,7 +32,7 @@ pub fn enter_ssa_impl<'a>(
32
33 let mut states = builder.complete();
34
35 - for block in fun.body.blocks.values_mut() {
35 + for block in fun.body.blocks.iter_mut() {
36 let state = states.remove(&block.id).unwrap();
37 block.phis = state.phis;
38 }
@@ -112,8 +112,9 @@ struct IncompletePhi<'a> {
112 impl<'a, 'e, 'f> Builder<'a, 'e, 'f> {
113 fn new(env: &'e Environment<'a>, entry: BlockId, blocks: &'f Blocks<'a>) -> Self {
114 let states = blocks
115 - .keys()
116 - .map(|block_id| (*block_id, BlockState::new(env)))
115 + .block_ids()
116 + .into_iter()
117 + .map(|block_id| (block_id, BlockState::new(env)))
118 .collect();
119 Self {
120 env,
@@ -182,7 +183,7 @@ impl<'a, 'e, 'f> Builder<'a, 'e, 'f> {
183 return identifier.clone();
184 }
185 // Else we have to look at predecessor blocks: bail if no predecessors
185 - let block = self.blocks.get(&block_id).unwrap();
186 + let block = self.blocks.block(block_id);
187 if block.predecessors.is_empty() {
188 println!("Unable to find previous id for {old_identifier:?}");
189 self.unknown.insert(old_identifier.id);
@@ -225,7 +226,7 @@ impl<'a, 'e, 'f> Builder<'a, 'e, 'f> {
226 identifier: new_identifier.clone(),
227 operands: Default::default(),
228 };
228 - let block = self.blocks.get(&block_id).unwrap();
229 + let block = self.blocks.block(block_id);
230 let preds = block.predecessors.clone();
231 for pred_block_id in preds {
232 let pred_id = self.get_id_at(pred_block_id, old_identifier);
@@ -262,15 +263,15 @@ impl<'a, 'e, 'f> Builder<'a, 'e, 'f> {
263 F: FnMut(&BasicBlock<'a>, &mut Self) -> Result<(), Diagnostic>,
264 {
265 let mut visited = IndexSet::new();
265 - let block_ids: Vec<_> = self.blocks.keys().cloned().collect_in(self.env.allocator);
266 + let block_ids = self.blocks.block_ids();
267 for block_id in block_ids {
268 visited.insert(block_id);
269 self.current = block_id;
269 - let block = self.blocks.get(&block_id).unwrap();
270 + let block = self.blocks.block(block_id);
271 f(block, self)?;
272 let successors = block.terminal.value.successors();
273 for successor in successors {
273 - let block = self.blocks.get(&successor).unwrap();
274 + let block = self.blocks.block(successor);
275 let count = self
276 .unsealed_predecessors
277 .entry(successor)