1use crate::nodes::base::{DspNode, ProcessContext};
2use dirtydata_core::types::{StableId, ConfigSnapshot};
3use std::collections::HashMap;
4
5#[derive(Clone, Debug)]
8pub enum DspOp {
9 LoadConst { val: f32, out: usize },
11 Copy { src: usize, dst: usize },
12
13 Add { a: usize, b: usize, out: usize },
15 Mul { a: usize, b: usize, out: usize },
16 Sin { src: usize, out: usize },
17
18 Accumulate { reg: usize, delta_reg: usize, wrap: f32 },
21
22 Tanh { src: usize, out: usize },
24
25 AssertRange { reg: usize, min: f32, max: f32, node_id: StableId },
29
30 CallLegacy { node_idx: usize, input_regs: Vec<usize>, output_regs: Vec<usize> },
33}
34
35pub struct JitProgram {
36 pub ops: Vec<DspOp>,
37 pub registers: Vec<[f32; 2]>,
38 pub legacy_nodes: Vec<Box<dyn DspNode>>,
39 pub constraint_violations: HashMap<StableId, String>,
41}
42
43impl JitProgram {
44 pub fn new() -> Self {
45 Self {
46 ops: Vec::new(),
47 registers: vec![[0.0; 2]; 1024], legacy_nodes: Vec::new(),
49 constraint_violations: HashMap::new(),
50 }
51 }
52
53 #[inline(always)]
54 pub fn execute(&mut self, ctx: &ProcessContext) -> [f32; 2] {
55 for op in &self.ops {
56 match op {
57 DspOp::LoadConst { val, out } => {
58 self.registers[*out] = [*val, *val];
59 }
60 DspOp::Copy { src, dst } => {
61 self.registers[*dst] = self.registers[*src];
62 }
63 DspOp::Add { a, b, out } => {
64 let v1 = self.registers[*a];
65 let v2 = self.registers[*b];
66 self.registers[*out] = [v1[0] + v2[0], v1[1] + v2[1]];
67 }
68 DspOp::Mul { a, b, out } => {
69 let v1 = self.registers[*a];
70 let v2 = self.registers[*b];
71 self.registers[*out] = [v1[0] * v2[0], v1[1] * v2[1]];
72 }
73 DspOp::Sin { src, out } => {
74 let v = self.registers[*src];
75 self.registers[*out] = [
76 (v[0] * 2.0 * std::f32::consts::PI).sin(),
77 (v[1] * 2.0 * std::f32::consts::PI).sin()
78 ];
79 }
80 DspOp::Accumulate { reg, delta_reg, wrap } => {
81 let mut v = self.registers[*reg];
82 let d = self.registers[*delta_reg];
83 for i in 0..2 {
84 v[i] = (v[i] + d[i]) % *wrap;
85 }
86 self.registers[*reg] = v;
87 }
88 DspOp::Tanh { src, out } => {
89 let v = self.registers[*src];
90 self.registers[*out] = [v[0].tanh(), v[1].tanh()];
91 }
92 DspOp::AssertRange { reg, min, max, node_id } => {
93 let v = self.registers[*reg];
94 if v[0] < *min || v[0] > *max || v[1] < *min || v[1] > *max {
95 if let Some(diag) = ctx.node_diagnostics {
96 diag.insert(*node_id, crate::DiagnosticRecord {
97 message: format!("Constraint Violation: Value {:.2} out of [{}, {}]", v[0], min, max),
98 severity: crate::DiagnosticSeverity::Warning,
99 timestamp: ctx.global_sample_index,
100 });
101 }
102 }
103 }
104 DspOp::CallLegacy { node_idx, input_regs, output_regs } => {
105 let node = &mut self.legacy_nodes[*node_idx];
106 let mut inputs = vec![0.0; input_regs.len()];
108 for (i, ®) in input_regs.iter().enumerate() {
109 inputs[i] = self.registers[reg][0];
110 }
111
112 let mut outputs = vec![[0.0; 2]; output_regs.len()];
113 node.process(&inputs, &mut outputs, &ConfigSnapshot::new(), ctx);
114
115 for (i, ®) in output_regs.iter().enumerate() {
116 self.registers[reg] = outputs[i];
117 }
118 }
119 }
120 }
121 self.registers[0] }
123}
124
125pub struct JitCompiler {
126 register_map: HashMap<StableId, usize>,
127 next_register: usize,
128 pub freeze_cache: HashMap<[u8; 32], std::path::PathBuf>,
130}
131
132impl JitCompiler {
133 pub fn new() -> Self {
134 Self {
135 register_map: HashMap::new(),
136 next_register: 1, freeze_cache: HashMap::new(),
138 }
139 }
140
141 pub fn compile_runner(&mut self, runner: &crate::DspRunner) -> JitProgram {
142 let mut program = JitProgram::new();
143 let graph = runner.get_graph();
144 let sample_rate = 44100.0;
145
146 for (id, _) in &runner.nodes {
148 self.register_map.insert(*id, self.next_register);
149 self.next_register += 1;
150 }
151
152 for (id, _node_impl) in &runner.nodes {
154 let out_reg = *self.register_map.get(id).unwrap();
155
156 if let Some(node_ir) = graph.nodes.get(id) {
157 match &node_ir.kind {
158 dirtydata_core::types::NodeKind::Source => {
159 let freq = node_ir.config.get("frequency").and_then(|v| v.as_float()).unwrap_or(440.0) as f32;
161
162 let delta_reg = self.next_register; self.next_register += 1;
163 program.ops.push(DspOp::LoadConst { val: freq / sample_rate, out: delta_reg });
164 program.ops.push(DspOp::Accumulate { reg: out_reg, delta_reg, wrap: 1.0 });
165 program.ops.push(DspOp::Sin { src: out_reg, out: out_reg });
166 }
167 dirtydata_core::types::NodeKind::Processor => {
168 let gain = node_ir.config.get("gain").and_then(|v| v.as_float()).unwrap_or(1.0) as f32;
170
171 let mut in_reg = 0;
172 for edge in graph.edges.values() {
173 if edge.target.node_id == *id {
174 if let Some(&src) = self.register_map.get(&edge.source.node_id) {
175 in_reg = src; break;
176 }
177 }
178 }
179
180 let gain_reg = self.next_register; self.next_register += 1;
181 program.ops.push(DspOp::LoadConst { val: gain, out: gain_reg });
182 program.ops.push(DspOp::Mul { a: in_reg, b: gain_reg, out: out_reg });
183 }
184 _ => {
185 program.ops.push(DspOp::CallLegacy {
186 node_idx: program.legacy_nodes.len(),
187 input_regs: vec![],
188 output_regs: vec![out_reg]
189 });
190 }
191 }
192
193 program.ops.push(DspOp::AssertRange {
195 reg: out_reg,
196 min: -2.0,
197 max: 2.0,
198 node_id: *id
199 });
200 }
201 }
202
203 if let Some(last_id) = runner.nodes.last().map(|(id, _)| id) {
204 let last_reg = *self.register_map.get(last_id).unwrap();
205 program.ops.push(DspOp::Copy { src: last_reg, dst: 0 });
206 }
207 program.ops.push(DspOp::Tanh { src: 0, out: 0 });
208
209 let optimizer = JitOptimizer::new();
211 optimizer.optimize(&mut program);
212
213 program
214 }
215}
216
217pub struct JitOptimizer {}
218
219impl JitOptimizer {
220 pub fn new() -> Self { Self {} }
221
222 pub fn optimize(&self, program: &mut JitProgram) {
223 self.common_subexpression_elimination(program);
224 self.constant_folding(program);
225 self.dead_code_elimination(program);
226 }
227
228 fn common_subexpression_elimination(&self, program: &mut JitProgram) {
229 #[derive(Hash, PartialEq, Eq)]
230 enum OpIdentity {
231 Sin { src: usize },
232 Add { a: usize, b: usize },
233 Mul { a: usize, b: usize },
234 Tanh { src: usize },
235 }
236
237 let mut available_expressions: HashMap<OpIdentity, usize> = HashMap::new();
238 let mut i = 0;
239 while i < program.ops.len() {
240 let identity = match &program.ops[i] {
241 DspOp::Sin { src, .. } => Some(OpIdentity::Sin { src: *src }),
242 DspOp::Add { a, b, .. } => Some(OpIdentity::Add { a: *a, b: *b }),
243 DspOp::Mul { a, b, .. } => Some(OpIdentity::Mul { a: *a, b: *b }),
244 DspOp::Tanh { src, .. } => Some(OpIdentity::Tanh { src: *src }),
245 _ => None,
246 };
247
248 if let Some(id) = identity {
249 if let Some(&prev_out) = available_expressions.get(&id) {
250 let current_out = match &program.ops[i] {
252 DspOp::Sin { out, .. } | DspOp::Add { out, .. } | DspOp::Mul { out, .. } | DspOp::Tanh { out, .. } => *out,
253 _ => unreachable!(),
254 };
255 program.ops[i] = DspOp::Copy { src: prev_out, dst: current_out };
256 } else {
257 let out = match &program.ops[i] {
258 DspOp::Sin { out, .. } | DspOp::Add { out, .. } | DspOp::Mul { out, .. } | DspOp::Tanh { out, .. } => *out,
259 _ => unreachable!(),
260 };
261 available_expressions.insert(id, out);
262 }
263 }
264 i += 1;
265 }
266 }
267
268 fn constant_folding(&self, program: &mut JitProgram) {
269 let mut constants: HashMap<usize, f32> = HashMap::new();
271 let mut i = 0;
272 while i < program.ops.len() {
273 let removed = false;
274 match &program.ops[i] {
275 DspOp::LoadConst { val, out } => {
276 constants.insert(*out, *val);
277 }
278 DspOp::Add { a, b, out } => {
279 if let (Some(&v1), Some(&v2)) = (constants.get(a), constants.get(b)) {
280 let result = v1 + v2;
281 constants.insert(*out, result);
282 program.ops[i] = DspOp::LoadConst { val: result, out: *out };
283 }
284 }
285 DspOp::Mul { a, b, out } => {
286 if let (Some(&v1), Some(&v2)) = (constants.get(a), constants.get(b)) {
287 let result = v1 * v2;
288 constants.insert(*out, result);
289 program.ops[i] = DspOp::LoadConst { val: result, out: *out };
290 }
291 }
292 _ => {
293 }
296 }
297 if !removed { i += 1; }
298 }
299 }
300
301 fn dead_code_elimination(&self, program: &mut JitProgram) {
302 let mut used_registers = std::collections::HashSet::new();
303 used_registers.insert(0); for op in program.ops.iter().rev() {
307 match op {
308 DspOp::Add { a, b, out } | DspOp::Mul { a, b, out } => {
309 if used_registers.contains(out) {
310 used_registers.insert(*a);
311 used_registers.insert(*b);
312 }
313 }
314 DspOp::Sin { src, out } | DspOp::Tanh { src, out } | DspOp::Copy { src, dst: out } => {
315 if used_registers.contains(out) {
316 used_registers.insert(*src);
317 }
318 }
319 DspOp::Accumulate { reg, delta_reg, .. } => {
320 used_registers.insert(*reg);
321 used_registers.insert(*delta_reg);
322 }
323 DspOp::AssertRange { reg, .. } => {
324 used_registers.insert(*reg); }
326 _ => {}
327 }
328 }
329
330 program.ops.retain(|op| {
332 match op {
333 DspOp::Add { out, .. } | DspOp::Mul { out, .. } | DspOp::Sin { out, .. } | DspOp::LoadConst { out, .. } => {
334 used_registers.contains(out)
335 }
336 _ => true }
338 });
339 }
340}