Skip to main content

shrew/exec/
jit.rs

1// =============================================================================
2// JIT Graph Compilation — Compile IR graphs into optimized execution plans
3// =============================================================================
4//
5// The default `Executor` interprets the IR graph on every call:
6//   - Recomputes topological order each time
7//   - Looks up each node in a HashMap
8//   - Matches on OpKind (50+ variants) per node
9//   - Allocates intermediates into a HashMap with no reuse
10//
11// The JIT compiler transforms the graph into a pre-compiled execution plan
12// that eliminates all of this overhead:
13//
14// COMPONENTS:
15//
16//   CompiledGraph    — The compiled execution plan for one IR graph
17//   Instruction      — A single operation in the compiled plan
18//   MemoryPlan       — Buffer lifecycle analysis and reuse
19//   BufferSlot       — A reusable memory slot (register analogy)
20//   JitExecutor      — Runs compiled graphs instead of re-interpreting
21//   CompileStats     — Compilation statistics
22//
23// WORKFLOW:
24//
25//   1. Compile:    JitExecutor::compile(program) → JitExecutor
26//   2. Run:        executor.run("Forward", &inputs) → JitResult
27//   3. Recompile:  executor.recompile("Forward") — after graph changes
28//
29// OPTIMIZATIONS:
30//
31//   - Pre-computed topological order (computed once at compile time)
32//   - Instruction tape (flat Vec<Instruction>, no HashMap lookup)
33//   - Memory planning: liveness analysis, buffer reuse (register allocation)
34//   - Dead value early-free (values dropped as soon as last consumer is done)
35//   - Fused dispatch (fused ops become single instructions)
36//   - Input/param index lookup pre-computed (no string matching at runtime)
37
38use std::collections::HashMap;
39use std::fmt;
40use std::time::Instant;
41
42use shrew_core::backend::Backend;
43use shrew_core::dtype::DType as CoreDType;
44use shrew_core::error::Result;
45use shrew_core::tensor::Tensor;
46
47use shrew_ir::graph::{ConstantValue, Dim, IrGraph, IrNode, IrProgram, IrType, NodeId, OpKind};
48
49use shrew_nn::Module;
50
51use super::engine::{ir_dtype_to_core, RuntimeConfig};
52
53// =============================================================================
54// Instruction — A single pre-compiled operation
55// =============================================================================
56
57/// Specifies which buffer slot an instruction reads from.
58#[derive(Debug, Clone)]
59pub struct SlotRef {
60    /// Index into the buffer table.
61    pub slot: usize,
62}
63
64/// A single operation in the compiled execution plan.
65///
66/// Unlike the interpreter, instructions are a flat enum with pre-resolved
67/// input/output buffer slots — no name lookups, no HashMap access.
68#[derive(Debug, Clone)]
69pub enum Instruction {
70    // ── Source instructions (produce values from external sources) ──
71    /// Load a graph input into a buffer slot.
72    LoadInput {
73        /// Input name (for lookup in the provided HashMap).
74        name: String,
75        /// Buffer slot to store the input tensor.
76        dst: usize,
77    },
78    /// Load a parameter into a buffer slot.
79    LoadParam {
80        /// Key: (graph_name, param_name).
81        graph_name: String,
82        param_name: String,
83        /// Buffer slot to store the parameter tensor.
84        dst: usize,
85    },
86
87    // ── Unary operations ──
88    Unary {
89        op: UnaryInstr,
90        src: usize,
91        dst: usize,
92    },
93
94    // ── Binary operations ──
95    Binary {
96        op: BinaryInstr,
97        lhs: usize,
98        rhs: usize,
99        dst: usize,
100    },
101
102    // ── Reduction operations ──
103    Reduce {
104        op: ReduceInstr,
105        src: usize,
106        dst: usize,
107        dims: Vec<i64>,
108        keepdim: bool,
109    },
110
111    // ── Shape operations ──
112    Reshape {
113        src: usize,
114        dst: usize,
115        /// Pre-resolved concrete shape (symbolic dims resolved at compile time).
116        shape: Vec<usize>,
117    },
118    Transpose {
119        src: usize,
120        dst: usize,
121    },
122    Permute {
123        src: usize,
124        dst: usize,
125        dims: Vec<i64>,
126    },
127    Expand {
128        src: usize,
129        dst: usize,
130        shape: Vec<usize>,
131    },
132    Concat {
133        srcs: Vec<usize>,
134        dst: usize,
135        dim: usize,
136    },
137    Split {
138        src: usize,
139        dst: usize,
140        dim: usize,
141        chunks: usize,
142    },
143
144    // ── Neural network operations ──
145    Softmax {
146        src: usize,
147        dst: usize,
148        dim: usize,
149    },
150    Embedding {
151        indices: usize,
152        table: usize,
153        dst: usize,
154    },
155    Linear {
156        input: usize,
157        weight: usize,
158        bias: Option<usize>,
159        dst: usize,
160    },
161    LayerNorm {
162        input: usize,
163        weight: usize,
164        bias: usize,
165        dst: usize,
166        eps: f64,
167    },
168    BatchNorm {
169        input: usize,
170        weight: Option<usize>,
171        bias: Option<usize>,
172        dst: usize,
173        eps: f64,
174    },
175    MultiHeadAttention {
176        input: usize,
177        dst: usize,
178        n_heads: usize,
179    },
180    TransformerBlock {
181        input: usize,
182        dst: usize,
183        n_heads: usize,
184    },
185    Dropout {
186        src: usize,
187        dst: usize,
188        p: f64,
189    },
190
191    // ── Loss functions ──
192    CrossEntropy {
193        predictions: usize,
194        targets: usize,
195        dst: usize,
196    },
197    MseLoss {
198        predictions: usize,
199        targets: usize,
200        dst: usize,
201    },
202
203    // ── Constants ──
204    Constant {
205        value: ConstantValue,
206        output_type: IrType,
207        dst: usize,
208    },
209
210    // ── Control flow ──
211    Repeat {
212        count: i64,
213        body_op: Box<OpKind>,
214        src: usize,
215        dst: usize,
216    },
217    Call {
218        graph_name: String,
219        inputs: Vec<usize>,
220        dst: usize,
221    },
222
223    // ── Comparison / logical ──
224    Compare {
225        op: CompareInstr,
226        lhs: usize,
227        rhs: usize,
228        dst: usize,
229    },
230    LogicalNot {
231        src: usize,
232        dst: usize,
233    },
234    LogicalBinOp {
235        op: LogicalBinInstr,
236        lhs: usize,
237        rhs: usize,
238        dst: usize,
239    },
240
241    // ── Fused operations (from IR optimizer) ──
242    FusedMatMulAdd {
243        a: usize,
244        b: usize,
245        c: usize,
246        dst: usize,
247    },
248    FusedAddRelu {
249        lhs: usize,
250        rhs: usize,
251        dst: usize,
252    },
253    FusedSubRelu {
254        lhs: usize,
255        rhs: usize,
256        dst: usize,
257    },
258    FusedMatMulRelu {
259        lhs: usize,
260        rhs: usize,
261        dst: usize,
262    },
263
264    // ── Identity (pass-through) ──
265    Copy {
266        src: usize,
267        dst: usize,
268    },
269
270    // ── Range ──
271    Range {
272        inputs: Vec<usize>,
273        output_type: IrType,
274        dst: usize,
275    },
276
277    // ── Free a buffer slot (dead value elimination) ──
278    Free {
279        slot: usize,
280    },
281}
282
283/// Unary operation variants (pre-dispatched).
284#[derive(Debug, Clone, Copy)]
285pub enum UnaryInstr {
286    Neg,
287    Relu,
288    Gelu,
289    Silu,
290    Sigmoid,
291    Tanh,
292    Exp,
293    Log,
294    Sqrt,
295}
296
297/// Binary operation variants (pre-dispatched).
298#[derive(Debug, Clone, Copy)]
299pub enum BinaryInstr {
300    Add,
301    Sub,
302    Mul,
303    Div,
304    MatMul,
305    Pow,
306    Mod,
307}
308
309/// Reduction operation variants.
310#[derive(Debug, Clone, Copy)]
311pub enum ReduceInstr {
312    Sum,
313    Mean,
314    Max,
315    Min,
316    Variance,
317}
318
319/// Comparison operation variants.
320#[derive(Debug, Clone, Copy)]
321pub enum CompareInstr {
322    Equal,
323    NotEqual,
324    Less,
325    Greater,
326    LessEqual,
327    GreaterEqual,
328}
329
330/// Logical binary operation variants.
331#[derive(Debug, Clone, Copy)]
332pub enum LogicalBinInstr {
333    And,
334    Or,
335}
336
337// =============================================================================
338// MemoryPlan — Buffer lifecycle analysis
339// =============================================================================
340
341/// Tracks when each value is first produced and last consumed.
342#[derive(Debug, Clone)]
343pub struct ValueLifetime {
344    /// Instruction index where this value is produced.
345    pub produced_at: usize,
346    /// Instruction index where this value is last consumed (inclusive).
347    pub last_used_at: usize,
348    /// Node ID from the original graph.
349    pub node_id: NodeId,
350    /// Whether this value is a graph output (must not be freed).
351    pub is_output: bool,
352    /// Whether this value is an input or parameter (externally owned).
353    pub is_external: bool,
354}
355
356/// The memory plan for a compiled graph — maps NodeIds to buffer slots
357/// and tracks lifetimes for dead value elimination.
358#[derive(Debug, Clone)]
359pub struct MemoryPlan {
360    /// Number of buffer slots needed.
361    pub num_slots: usize,
362    /// Mapping from NodeId → buffer slot.
363    pub node_to_slot: HashMap<usize, usize>,
364    /// Lifetime of each slot.
365    pub lifetimes: Vec<ValueLifetime>,
366    /// Free instructions to insert (slot, after_instruction_idx).
367    pub free_points: Vec<(usize, usize)>,
368    /// Number of buffers reused.
369    pub reuse_count: usize,
370}
371
372// =============================================================================
373// CompiledGraph — A fully compiled execution plan
374// =============================================================================
375
376/// The compiled execution plan for a single IR graph.
377///
378/// Contains a flat instruction tape, memory plan, and metadata for
379/// efficient repeated execution.
380#[derive(Debug)]
381pub struct CompiledGraph {
382    /// Name of the source graph.
383    pub graph_name: String,
384    /// Flat instruction tape — executed sequentially.
385    pub instructions: Vec<Instruction>,
386    /// Memory plan — buffer slot assignments.
387    pub memory_plan: MemoryPlan,
388    /// Output slot mappings: name → buffer slot.
389    pub output_slots: HashMap<String, usize>,
390    /// Total number of buffer slots.
391    pub num_slots: usize,
392    /// Compilation statistics.
393    pub stats: CompileStats,
394}
395
396/// Statistics from the compilation process.
397#[derive(Debug, Clone)]
398pub struct CompileStats {
399    /// Number of instructions in the compiled plan.
400    pub num_instructions: usize,
401    /// Number of nodes in the source graph.
402    pub num_source_nodes: usize,
403    /// Number of buffer slots allocated.
404    pub num_slots: usize,
405    /// Number of buffer slots reused.
406    pub num_reused: usize,
407    /// Number of free instructions inserted.
408    pub num_frees: usize,
409    /// Number of fused instructions.
410    pub num_fused: usize,
411    /// Compilation time in microseconds.
412    pub compile_time_us: u64,
413}
414
415impl fmt::Display for CompileStats {
416    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
417        write!(
418            f,
419            "CompiledGraph: {} instructions ({} source nodes), {} slots ({} reused), {} frees, {} fused, compiled in {}μs",
420            self.num_instructions,
421            self.num_source_nodes,
422            self.num_slots,
423            self.num_reused,
424            self.num_frees,
425            self.num_fused,
426            self.compile_time_us,
427        )
428    }
429}
430
431// =============================================================================
432// Compilation — transform IrGraph → CompiledGraph
433// =============================================================================
434
435/// Compile a single IR graph into an optimized execution plan.
436pub fn compile_graph(
437    graph: &IrGraph,
438    program: &IrProgram,
439    config: &RuntimeConfig,
440) -> Result<CompiledGraph> {
441    let start = Instant::now();
442
443    // 1. Get topological order (computed once)
444    let order = graph.topo_order();
445
446    // 2. Assign buffer slots
447    let mut node_to_slot: HashMap<usize, usize> = HashMap::new();
448    let mut next_slot = 0usize;
449    let mut produced_at: HashMap<usize, usize> = HashMap::new();
450    let mut last_used_at: HashMap<usize, usize> = HashMap::new();
451
452    // Pre-compute output node IDs for quick lookup
453    let output_node_ids: std::collections::HashSet<usize> =
454        graph.outputs.iter().map(|o| o.node_id.0).collect();
455
456    // Pre-compute input/param node IDs
457    let input_node_ids: std::collections::HashSet<usize> =
458        graph.inputs.iter().map(|id| id.0).collect();
459    let param_node_ids: std::collections::HashSet<usize> =
460        graph.params.iter().map(|p| p.node_id.0).collect();
461
462    // Assign a slot to each node in topo order
463    for &node_id in &order {
464        let slot = next_slot;
465        node_to_slot.insert(node_id.0, slot);
466        next_slot += 1;
467    }
468
469    // 3. Build instruction tape
470    let mut instructions: Vec<Instruction> = Vec::with_capacity(order.len() + graph.inputs.len());
471    let mut num_fused = 0;
472
473    for (instr_idx, &node_id) in order.iter().enumerate() {
474        let node = graph.node(node_id);
475        let dst = node_to_slot[&node_id.0];
476
477        // Track production point
478        produced_at.insert(node_id.0, instr_idx);
479
480        // Track consumption points for inputs
481        for &input_id in &node.inputs {
482            last_used_at.insert(input_id.0, instr_idx);
483        }
484
485        // Check if this node is an input
486        if input_node_ids.contains(&node_id.0) {
487            instructions.push(Instruction::LoadInput {
488                name: node.name.clone(),
489                dst,
490            });
491            continue;
492        }
493
494        // Check if this node is a parameter
495        if param_node_ids.contains(&node_id.0) {
496            if let Some(param) = graph.params.iter().find(|p| p.node_id == node_id) {
497                instructions.push(Instruction::LoadParam {
498                    graph_name: graph.name.clone(),
499                    param_name: param.name.clone(),
500                    dst,
501                });
502            }
503            continue;
504        }
505
506        // Compile the operation
507        let instr = compile_node(graph, node, &node_to_slot, config, program)?;
508
509        // Track fused instructions
510        match &instr {
511            Instruction::FusedMatMulAdd { .. }
512            | Instruction::FusedAddRelu { .. }
513            | Instruction::FusedSubRelu { .. }
514            | Instruction::FusedMatMulRelu { .. } => {
515                num_fused += 1;
516            }
517            _ => {}
518        }
519
520        instructions.push(instr);
521    }
522
523    // 4. Compute lifetimes and insert free instructions
524    let mut lifetimes = Vec::new();
525    let mut free_points = Vec::new();
526
527    for &node_id in &order {
528        let slot = node_to_slot[&node_id.0];
529        let is_output = output_node_ids.contains(&node_id.0);
530        let is_external =
531            input_node_ids.contains(&node_id.0) || param_node_ids.contains(&node_id.0);
532        let prod = produced_at.get(&node_id.0).copied().unwrap_or(0);
533        let last = last_used_at.get(&node_id.0).copied().unwrap_or(prod);
534
535        lifetimes.push(ValueLifetime {
536            produced_at: prod,
537            last_used_at: last,
538            node_id,
539            is_output,
540            is_external,
541        });
542
543        // Insert free point if this value is not an output and not external
544        if !is_output && !is_external && last < instructions.len().saturating_sub(1) {
545            free_points.push((slot, last));
546        }
547    }
548
549    // Sort frees by position (latest first for stable insertion)
550    free_points.sort_by(|a, b| b.1.cmp(&a.1));
551
552    // Insert free instructions after the last use
553    let num_frees = free_points.len();
554    for (slot, after_idx) in &free_points {
555        let insert_pos = (*after_idx + 1).min(instructions.len());
556        instructions.insert(insert_pos, Instruction::Free { slot: *slot });
557    }
558
559    // 5. Build output slot mapping
560    let mut output_slots = HashMap::new();
561    for output in &graph.outputs {
562        if let Some(&slot) = node_to_slot.get(&output.node_id.0) {
563            output_slots.insert(output.name.clone(), slot);
564        }
565    }
566
567    let num_slots = next_slot;
568    let compile_time = start.elapsed();
569
570    let memory_plan = MemoryPlan {
571        num_slots,
572        node_to_slot,
573        lifetimes,
574        free_points: Vec::new(), // Already applied
575        reuse_count: 0,          // No physical reuse in this version (logical slots)
576    };
577
578    let stats = CompileStats {
579        num_instructions: instructions.len(),
580        num_source_nodes: graph.nodes.len(),
581        num_slots,
582        num_reused: 0,
583        num_frees,
584        num_fused,
585        compile_time_us: compile_time.as_micros() as u64,
586    };
587
588    Ok(CompiledGraph {
589        graph_name: graph.name.clone(),
590        instructions,
591        memory_plan,
592        output_slots,
593        num_slots,
594        stats,
595    })
596}
597
598/// Compile a single IR node into an instruction.
599fn compile_node(
600    _graph: &IrGraph,
601    node: &IrNode,
602    node_to_slot: &HashMap<usize, usize>,
603    config: &RuntimeConfig,
604    program: &IrProgram,
605) -> Result<Instruction> {
606    let dst = node_to_slot[&node.id.0];
607
608    // Helper: get slot for an input
609    let slot = |idx: usize| -> Result<usize> {
610        let input_id = node.inputs.get(idx).ok_or_else(|| {
611            shrew_core::Error::msg(format!(
612                "Node '{}' expected input at index {}, but has {} inputs",
613                node.name,
614                idx,
615                node.inputs.len()
616            ))
617        })?;
618        node_to_slot.get(&input_id.0).copied().ok_or_else(|| {
619            shrew_core::Error::msg(format!(
620                "Node '{}' input {} (NodeId {}) not found in slot map",
621                node.name, idx, input_id.0
622            ))
623        })
624    };
625
626    match &node.op {
627        OpKind::Identity => Ok(Instruction::Copy { src: slot(0)?, dst }),
628
629        // ── Unary ──
630        OpKind::Neg => Ok(Instruction::Unary {
631            op: UnaryInstr::Neg,
632            src: slot(0)?,
633            dst,
634        }),
635        OpKind::Relu => Ok(Instruction::Unary {
636            op: UnaryInstr::Relu,
637            src: slot(0)?,
638            dst,
639        }),
640        OpKind::Gelu => Ok(Instruction::Unary {
641            op: UnaryInstr::Gelu,
642            src: slot(0)?,
643            dst,
644        }),
645        OpKind::Silu => Ok(Instruction::Unary {
646            op: UnaryInstr::Silu,
647            src: slot(0)?,
648            dst,
649        }),
650        OpKind::Sigmoid => Ok(Instruction::Unary {
651            op: UnaryInstr::Sigmoid,
652            src: slot(0)?,
653            dst,
654        }),
655        OpKind::Tanh => Ok(Instruction::Unary {
656            op: UnaryInstr::Tanh,
657            src: slot(0)?,
658            dst,
659        }),
660        OpKind::Exp => Ok(Instruction::Unary {
661            op: UnaryInstr::Exp,
662            src: slot(0)?,
663            dst,
664        }),
665        OpKind::Log => Ok(Instruction::Unary {
666            op: UnaryInstr::Log,
667            src: slot(0)?,
668            dst,
669        }),
670        OpKind::Sqrt => Ok(Instruction::Unary {
671            op: UnaryInstr::Sqrt,
672            src: slot(0)?,
673            dst,
674        }),
675
676        // ── Binary ──
677        OpKind::Add => Ok(Instruction::Binary {
678            op: BinaryInstr::Add,
679            lhs: slot(0)?,
680            rhs: slot(1)?,
681            dst,
682        }),
683        OpKind::Sub => Ok(Instruction::Binary {
684            op: BinaryInstr::Sub,
685            lhs: slot(0)?,
686            rhs: slot(1)?,
687            dst,
688        }),
689        OpKind::Mul => Ok(Instruction::Binary {
690            op: BinaryInstr::Mul,
691            lhs: slot(0)?,
692            rhs: slot(1)?,
693            dst,
694        }),
695        OpKind::Div => Ok(Instruction::Binary {
696            op: BinaryInstr::Div,
697            lhs: slot(0)?,
698            rhs: slot(1)?,
699            dst,
700        }),
701        OpKind::MatMul => Ok(Instruction::Binary {
702            op: BinaryInstr::MatMul,
703            lhs: slot(0)?,
704            rhs: slot(1)?,
705            dst,
706        }),
707        OpKind::Pow => Ok(Instruction::Binary {
708            op: BinaryInstr::Pow,
709            lhs: slot(0)?,
710            rhs: slot(1)?,
711            dst,
712        }),
713        OpKind::Mod => Ok(Instruction::Binary {
714            op: BinaryInstr::Mod,
715            lhs: slot(0)?,
716            rhs: slot(1)?,
717            dst,
718        }),
719
720        // ── Transpose ──
721        OpKind::Transpose => Ok(Instruction::Transpose { src: slot(0)?, dst }),
722
723        // ── Reductions ──
724        OpKind::Sum { dims, keepdim } => Ok(Instruction::Reduce {
725            op: ReduceInstr::Sum,
726            src: slot(0)?,
727            dst,
728            dims: dims.clone(),
729            keepdim: *keepdim,
730        }),
731        OpKind::Mean { dims, keepdim } => Ok(Instruction::Reduce {
732            op: ReduceInstr::Mean,
733            src: slot(0)?,
734            dst,
735            dims: dims.clone(),
736            keepdim: *keepdim,
737        }),
738        OpKind::Max { dim, keepdim } => Ok(Instruction::Reduce {
739            op: ReduceInstr::Max,
740            src: slot(0)?,
741            dst,
742            dims: vec![*dim],
743            keepdim: *keepdim,
744        }),
745        OpKind::Min { dim, keepdim } => Ok(Instruction::Reduce {
746            op: ReduceInstr::Min,
747            src: slot(0)?,
748            dst,
749            dims: vec![*dim],
750            keepdim: *keepdim,
751        }),
752        OpKind::Variance { dims, keepdim } => Ok(Instruction::Reduce {
753            op: ReduceInstr::Variance,
754            src: slot(0)?,
755            dst,
756            dims: dims.clone(),
757            keepdim: *keepdim,
758        }),
759
760        // ── Softmax ──
761        OpKind::Softmax { dim } => {
762            let d = *dim;
763            Ok(Instruction::Softmax {
764                src: slot(0)?,
765                dst,
766                dim: d as usize,
767            })
768        }
769
770        // ── Shape ops ──
771        OpKind::Reshape { target_shape } | OpKind::View { target_shape } => {
772            let shape = resolve_shape_vec(target_shape, config, program)?;
773            Ok(Instruction::Reshape {
774                src: slot(0)?,
775                dst,
776                shape,
777            })
778        }
779        OpKind::Permute { dims } => Ok(Instruction::Permute {
780            src: slot(0)?,
781            dst,
782            dims: dims.clone(),
783        }),
784        OpKind::Expand { target_shape } => {
785            let shape = resolve_shape_vec(target_shape, config, program)?;
786            Ok(Instruction::Expand {
787                src: slot(0)?,
788                dst,
789                shape,
790            })
791        }
792        OpKind::Concat { dim } => {
793            let srcs: Vec<usize> = (0..node.inputs.len())
794                .map(&slot)
795                .collect::<Result<Vec<_>>>()?;
796            Ok(Instruction::Concat {
797                srcs,
798                dst,
799                dim: *dim as usize,
800            })
801        }
802        OpKind::Split { dim, chunks } => Ok(Instruction::Split {
803            src: slot(0)?,
804            dst,
805            dim: resolve_neg_dim(*dim, 4), // dim resolved at runtime
806            chunks: *chunks as usize,
807        }),
808
809        // ── NN layers ──
810        OpKind::Embedding => Ok(Instruction::Embedding {
811            indices: slot(0)?,
812            table: slot(1)?,
813            dst,
814        }),
815        OpKind::Linear { bias } => {
816            let bias_slot = if *bias && node.inputs.len() >= 3 {
817                Some(slot(2)?)
818            } else {
819                None
820            };
821            Ok(Instruction::Linear {
822                input: slot(0)?,
823                weight: slot(1)?,
824                bias: bias_slot,
825                dst,
826            })
827        }
828        OpKind::LayerNorm { eps } => Ok(Instruction::LayerNorm {
829            input: slot(0)?,
830            weight: slot(1)?,
831            bias: slot(2)?,
832            dst,
833            eps: *eps,
834        }),
835        OpKind::BatchNorm { eps } => {
836            let weight = if node.inputs.len() >= 2 {
837                Some(slot(1)?)
838            } else {
839                None
840            };
841            let bias = if node.inputs.len() >= 3 {
842                Some(slot(2)?)
843            } else {
844                None
845            };
846            Ok(Instruction::BatchNorm {
847                input: slot(0)?,
848                weight,
849                bias,
850                dst,
851                eps: *eps,
852            })
853        }
854        OpKind::MultiHeadAttention { n_heads } => Ok(Instruction::MultiHeadAttention {
855            input: slot(0)?,
856            dst,
857            n_heads: *n_heads as usize,
858        }),
859        OpKind::TransformerBlock { n_heads } => Ok(Instruction::TransformerBlock {
860            input: slot(0)?,
861            dst,
862            n_heads: *n_heads as usize,
863        }),
864        OpKind::Dropout { p } => Ok(Instruction::Dropout {
865            src: slot(0)?,
866            dst,
867            p: *p,
868        }),
869
870        // ── Loss ──
871        OpKind::CrossEntropy => Ok(Instruction::CrossEntropy {
872            predictions: slot(0)?,
873            targets: slot(1)?,
874            dst,
875        }),
876        OpKind::MseLoss => Ok(Instruction::MseLoss {
877            predictions: slot(0)?,
878            targets: slot(1)?,
879            dst,
880        }),
881
882        // ── Constants ──
883        OpKind::Constant(val) => Ok(Instruction::Constant {
884            value: val.clone(),
885            output_type: node.output_type.clone(),
886            dst,
887        }),
888
889        // ── Repeat ──
890        OpKind::Repeat { count, body_op } => Ok(Instruction::Repeat {
891            count: *count,
892            body_op: body_op.clone(),
893            src: slot(0)?,
894            dst,
895        }),
896
897        // ── Call ──
898        OpKind::Call { graph_name } => {
899            let inputs: Vec<usize> = (0..node.inputs.len())
900                .map(&slot)
901                .collect::<Result<Vec<_>>>()?;
902            Ok(Instruction::Call {
903                graph_name: graph_name.clone(),
904                inputs,
905                dst,
906            })
907        }
908
909        // ── Range ──
910        OpKind::Range => {
911            let inputs: Vec<usize> = (0..node.inputs.len())
912                .map(&slot)
913                .collect::<Result<Vec<_>>>()?;
914            Ok(Instruction::Range {
915                inputs,
916                output_type: node.output_type.clone(),
917                dst,
918            })
919        }
920
921        // ── Comparison ──
922        OpKind::Equal => Ok(Instruction::Compare {
923            op: CompareInstr::Equal,
924            lhs: slot(0)?,
925            rhs: slot(1)?,
926            dst,
927        }),
928        OpKind::NotEqual => Ok(Instruction::Compare {
929            op: CompareInstr::NotEqual,
930            lhs: slot(0)?,
931            rhs: slot(1)?,
932            dst,
933        }),
934        OpKind::Less => Ok(Instruction::Compare {
935            op: CompareInstr::Less,
936            lhs: slot(0)?,
937            rhs: slot(1)?,
938            dst,
939        }),
940        OpKind::Greater => Ok(Instruction::Compare {
941            op: CompareInstr::Greater,
942            lhs: slot(0)?,
943            rhs: slot(1)?,
944            dst,
945        }),
946        OpKind::LessEqual => Ok(Instruction::Compare {
947            op: CompareInstr::LessEqual,
948            lhs: slot(0)?,
949            rhs: slot(1)?,
950            dst,
951        }),
952        OpKind::GreaterEqual => Ok(Instruction::Compare {
953            op: CompareInstr::GreaterEqual,
954            lhs: slot(0)?,
955            rhs: slot(1)?,
956            dst,
957        }),
958
959        // ── Logical ──
960        OpKind::And => Ok(Instruction::LogicalBinOp {
961            op: LogicalBinInstr::And,
962            lhs: slot(0)?,
963            rhs: slot(1)?,
964            dst,
965        }),
966        OpKind::Or => Ok(Instruction::LogicalBinOp {
967            op: LogicalBinInstr::Or,
968            lhs: slot(0)?,
969            rhs: slot(1)?,
970            dst,
971        }),
972        OpKind::Not => Ok(Instruction::LogicalNot { src: slot(0)?, dst }),
973
974        // ── Fused ops (from IR optimizer) ──
975        OpKind::Custom { name, .. } => match name.as_str() {
976            "fused_matmul_add" => Ok(Instruction::FusedMatMulAdd {
977                a: slot(0)?,
978                b: slot(1)?,
979                c: slot(2)?,
980                dst,
981            }),
982            "fused_add_relu" => Ok(Instruction::FusedAddRelu {
983                lhs: slot(0)?,
984                rhs: slot(1)?,
985                dst,
986            }),
987            "fused_sub_relu" => Ok(Instruction::FusedSubRelu {
988                lhs: slot(0)?,
989                rhs: slot(1)?,
990                dst,
991            }),
992            "fused_matmul_relu" => Ok(Instruction::FusedMatMulRelu {
993                lhs: slot(0)?,
994                rhs: slot(1)?,
995                dst,
996            }),
997            other => Err(shrew_core::Error::msg(format!(
998                "Unknown custom op '{}' during JIT compilation",
999                other
1000            ))),
1001        },
1002    }
1003}
1004
1005// =============================================================================
1006// JitExecutor — Runs compiled graphs
1007// =============================================================================
1008
1009/// A JIT-compiled executor that runs pre-compiled graph execution plans.
1010///
1011/// Unlike the interpreter (`Executor`), the JIT executor:
1012/// - Pre-compiles each graph into a flat instruction tape
1013/// - Pre-resolves all buffer slot assignments
1014/// - Inserts dead-value-free instructions for memory efficiency
1015/// - Dispatches operations without HashMap lookups or string matching
1016///
1017/// # Usage
1018/// ```ignore
1019/// let jit = JitExecutor::<CpuBackend>::compile(program, device, config)?;
1020/// let result = jit.run("Forward", &inputs)?;
1021/// let output = result.get("output").unwrap();
1022/// ```
1023pub struct JitExecutor<B: Backend> {
1024    /// Compiled graphs, keyed by graph name.
1025    compiled: HashMap<String, CompiledGraph>,
1026    /// The source IR program.
1027    program: IrProgram,
1028    /// Runtime configuration.
1029    config: RuntimeConfig,
1030    /// Device.
1031    device: B::Device,
1032    /// Initialized parameters.
1033    params: HashMap<(String, String), Tensor<B>>,
1034    /// Persistent transformer blocks (keyed by dst slot).
1035    transformer_blocks: std::cell::RefCell<HashMap<usize, shrew_nn::TransformerBlock<B>>>,
1036    /// Persistent multi-head attention blocks (keyed by dst slot).
1037    mha_blocks: std::cell::RefCell<HashMap<usize, shrew_nn::MultiHeadAttention<B>>>,
1038    /// Persistent blocks inside Repeat instructions (keyed by (dst slot, repeat index)).
1039    repeat_blocks: std::cell::RefCell<HashMap<(usize, usize), shrew_nn::TransformerBlock<B>>>,
1040}
1041
1042/// Result of a JIT execution.
1043#[derive(Debug)]
1044pub struct JitResult<B: Backend> {
1045    /// Output tensors, keyed by output name.
1046    pub outputs: HashMap<String, Tensor<B>>,
1047}
1048
1049impl<B: Backend> JitResult<B> {
1050    /// Get the first output tensor.
1051    pub fn output(&self) -> Option<&Tensor<B>> {
1052        self.outputs.values().next()
1053    }
1054
1055    /// Get an output by name.
1056    pub fn get(&self, name: &str) -> Option<&Tensor<B>> {
1057        self.outputs.get(name)
1058    }
1059}
1060
1061impl<B: Backend> JitExecutor<B> {
1062    /// Compile all graphs in a program and create a JIT executor.
1063    pub fn compile(program: IrProgram, device: B::Device, config: RuntimeConfig) -> Result<Self> {
1064        let mut compiled = HashMap::new();
1065
1066        // Compile each graph
1067        for graph in &program.graphs {
1068            let cg = compile_graph(graph, &program, &config)?;
1069            compiled.insert(graph.name.clone(), cg);
1070        }
1071
1072        // Initialize parameters (reuse logic from Executor)
1073        let mut params = HashMap::new();
1074        for graph in &program.graphs {
1075            for param in &graph.params {
1076                let tensor = init_param::<B>(
1077                    &param.ty,
1078                    &param.init,
1079                    param.frozen,
1080                    &config,
1081                    &program,
1082                    &device,
1083                )?;
1084                params.insert((graph.name.clone(), param.name.clone()), tensor);
1085            }
1086        }
1087
1088        Ok(Self {
1089            compiled,
1090            program,
1091            config,
1092            device,
1093            params,
1094            transformer_blocks: std::cell::RefCell::new(HashMap::new()),
1095            mha_blocks: std::cell::RefCell::new(HashMap::new()),
1096            repeat_blocks: std::cell::RefCell::new(HashMap::new()),
1097        })
1098    }
1099
1100
1101    /// Get compilation statistics for a graph.
1102    pub fn stats(&self, graph_name: &str) -> Option<&CompileStats> {
1103        self.compiled.get(graph_name).map(|cg| &cg.stats)
1104    }
1105
1106    /// Get all compilation statistics.
1107    pub fn all_stats(&self) -> Vec<(&str, &CompileStats)> {
1108        self.compiled
1109            .iter()
1110            .map(|(name, cg)| (name.as_str(), &cg.stats))
1111            .collect()
1112    }
1113
1114    /// Run a compiled graph with the given inputs.
1115    pub fn run(
1116        &self,
1117        graph_name: &str,
1118        inputs: &HashMap<String, Tensor<B>>,
1119    ) -> Result<JitResult<B>> {
1120        let cg = self.compiled.get(graph_name).ok_or_else(|| {
1121            shrew_core::Error::msg(format!(
1122                "Graph '{}' not compiled. Available: {:?}",
1123                graph_name,
1124                self.compiled.keys().collect::<Vec<_>>()
1125            ))
1126        })?;
1127
1128        // Allocate buffer table (slots)
1129        let mut slots: Vec<Option<Tensor<B>>> = vec![None; cg.num_slots];
1130
1131        // Execute instruction tape
1132        for instr in &cg.instructions {
1133            match instr {
1134                Instruction::LoadInput { name, dst } => {
1135                    if let Some(tensor) = inputs.get(name) {
1136                        slots[*dst] = Some(tensor.clone());
1137                    }
1138                }
1139
1140                Instruction::LoadParam {
1141                    graph_name,
1142                    param_name,
1143                    dst,
1144                } => {
1145                    let key = (graph_name.clone(), param_name.clone());
1146                    if let Some(tensor) = self.params.get(&key) {
1147                        slots[*dst] = Some(tensor.clone());
1148                    }
1149                }
1150
1151                Instruction::Unary { op, src, dst } => {
1152                    let t = get_slot(&slots, *src)?;
1153                    let result = match op {
1154                        UnaryInstr::Neg => t.neg(),
1155                        UnaryInstr::Relu => t.relu(),
1156                        UnaryInstr::Gelu => t.gelu(),
1157                        UnaryInstr::Silu => t.silu(),
1158                        UnaryInstr::Sigmoid => t.sigmoid(),
1159                        UnaryInstr::Tanh => t.tanh(),
1160                        UnaryInstr::Exp => t.exp(),
1161                        UnaryInstr::Log => t.log(),
1162                        UnaryInstr::Sqrt => t.sqrt(),
1163                    }?;
1164                    slots[*dst] = Some(result);
1165                }
1166
1167                Instruction::Binary { op, lhs, rhs, dst } => {
1168                    let a = get_slot(&slots, *lhs)?;
1169                    let b = get_slot(&slots, *rhs)?;
1170                    let result = match op {
1171                        BinaryInstr::Add => a.add(b),
1172                        BinaryInstr::Sub => a.sub(b),
1173                        BinaryInstr::Mul => a.mul(b),
1174                        BinaryInstr::Div => a.div(b),
1175                        BinaryInstr::MatMul => a.matmul(b),
1176                        BinaryInstr::Pow => a.log()?.mul(b)?.exp(), // x^y = exp(y*ln(x))
1177                        BinaryInstr::Mod => {
1178                            let quotient = a.div(b)?.floor()?;
1179                            let product = quotient.mul(b)?;
1180                            a.sub(&product)
1181                        }
1182                    }?;
1183                    slots[*dst] = Some(result);
1184                }
1185
1186                Instruction::Reduce {
1187                    op,
1188                    src,
1189                    dst,
1190                    dims,
1191                    keepdim,
1192                } => {
1193                    let t = get_slot(&slots, *src)?;
1194                    let result = match op {
1195                        ReduceInstr::Sum => {
1196                            if dims.is_empty() || (dims.len() == 1 && dims[0] == -1) {
1197                                t.sum_all()
1198                            } else {
1199                                let d = resolve_neg_dim(dims[0], t.rank());
1200                                t.sum(d as usize, *keepdim)
1201                            }
1202                        }
1203                        ReduceInstr::Mean => {
1204                            if dims.is_empty() || (dims.len() == 1 && dims[0] == -1) {
1205                                t.mean_all()
1206                            } else {
1207                                let d = resolve_neg_dim(dims[0], t.rank());
1208                                t.mean(d as usize, *keepdim)
1209                            }
1210                        }
1211                        ReduceInstr::Max => {
1212                            let d = resolve_neg_dim(dims[0], t.rank());
1213                            t.max(d as usize, *keepdim)
1214                        }
1215                        ReduceInstr::Min => {
1216                            let d = resolve_neg_dim(dims[0], t.rank());
1217                            t.min(d as usize, *keepdim)
1218                        }
1219                        ReduceInstr::Variance => {
1220                            if dims.is_empty() {
1221                                t.var(0, *keepdim)
1222                            } else {
1223                                let d = resolve_neg_dim(dims[0], t.rank());
1224                                t.var(d as usize, *keepdim)
1225                            }
1226                        }
1227                    }?;
1228                    slots[*dst] = Some(result);
1229                }
1230
1231                Instruction::Reshape { src, dst, shape } => {
1232                    let t = get_slot(&slots, *src)?;
1233                    slots[*dst] = Some(t.reshape(shape.clone())?);
1234                }
1235
1236                Instruction::Transpose { src, dst } => {
1237                    let t = get_slot(&slots, *src)?;
1238                    let rank = t.rank();
1239                    slots[*dst] = Some(t.transpose(rank - 2, rank - 1)?);
1240                }
1241
1242                Instruction::Permute { src, dst, dims } => {
1243                    let t = get_slot(&slots, *src)?;
1244                    let mut result = t.clone();
1245                    let mut current: Vec<usize> = (0..t.rank()).collect();
1246                    for i in 0..dims.len() {
1247                        let target = dims[i] as usize;
1248                        if current[i] != target {
1249                            if let Some(j) = current.iter().position(|&x| x == target) {
1250                                result = result.transpose(i, j)?;
1251                                current.swap(i, j);
1252                            }
1253                        }
1254                    }
1255                    slots[*dst] = Some(result);
1256                }
1257
1258                Instruction::Expand { src, dst, shape } => {
1259                    let t = get_slot(&slots, *src)?;
1260                    slots[*dst] = Some(t.expand(shape.clone())?);
1261                }
1262
1263                Instruction::Concat { srcs, dst, dim } => {
1264                    let tensors: Vec<Tensor<B>> = srcs
1265                        .iter()
1266                        .map(|s| get_slot(&slots, *s).cloned())
1267                        .collect::<Result<Vec<_>>>()?;
1268                    slots[*dst] = Some(Tensor::<B>::cat(&tensors, *dim)?);
1269                }
1270
1271                Instruction::Split {
1272                    src,
1273                    dst,
1274                    dim,
1275                    chunks,
1276                } => {
1277                    let t = get_slot(&slots, *src)?;
1278                    let result = t.chunk(*chunks, *dim)?;
1279                    if let Some(first) = result.into_iter().next() {
1280                        slots[*dst] = Some(first);
1281                    }
1282                }
1283
1284                Instruction::Softmax { src, dst, dim } => {
1285                    let t = get_slot(&slots, *src)?;
1286                    slots[*dst] = Some(t.softmax(*dim)?);
1287                }
1288
1289                Instruction::Embedding {
1290                    indices,
1291                    table,
1292                    dst,
1293                } => {
1294                    let idx = get_slot(&slots, *indices)?;
1295                    let tbl = get_slot(&slots, *table)?;
1296                    let emb = shrew_nn::Embedding::<B>::from_tensor(tbl.clone())?;
1297                    slots[*dst] = Some(emb.forward(idx)?);
1298                }
1299
1300                Instruction::Linear {
1301                    input,
1302                    weight,
1303                    bias,
1304                    dst,
1305                } => {
1306                    let inp = get_slot(&slots, *input)?;
1307                    let w = get_slot(&slots, *weight)?;
1308                    let b = bias.map(|s| get_slot(&slots, s).cloned()).transpose()?;
1309                    let lin = shrew_nn::Linear::<B>::from_tensors(w.clone(), b)?;
1310                    slots[*dst] = Some(lin.forward(inp)?);
1311                }
1312
1313                Instruction::LayerNorm {
1314                    input,
1315                    weight,
1316                    bias,
1317                    dst,
1318                    eps,
1319                } => {
1320                    let inp = get_slot(&slots, *input)?;
1321                    let w = get_slot(&slots, *weight)?;
1322                    let b = get_slot(&slots, *bias)?;
1323                    let ln = shrew_nn::LayerNorm::<B>::from_tensors(w.clone(), b.clone(), *eps)?;
1324                    slots[*dst] = Some(ln.forward(inp)?);
1325                }
1326
1327                Instruction::BatchNorm {
1328                    input,
1329                    weight,
1330                    bias,
1331                    dst,
1332                    eps,
1333                } => {
1334                    let inp = get_slot(&slots, *input)?;
1335                    if let (Some(ws), Some(bs)) = (weight, bias) {
1336                        let w = get_slot(&slots, *ws)?;
1337                        let b = get_slot(&slots, *bs)?;
1338                        let bn =
1339                            shrew_nn::BatchNorm2d::<B>::from_tensors(w.clone(), b.clone(), *eps)?;
1340                        slots[*dst] = Some(bn.forward(inp)?);
1341                    } else {
1342                        let dims = inp.dims();
1343                        let c = if dims.len() == 4 { dims[1] } else { dims[0] };
1344                        let bn = shrew_nn::BatchNorm2d::<B>::new(
1345                            c,
1346                            *eps,
1347                            0.1,
1348                            inp.dtype(),
1349                            &self.device,
1350                        )?;
1351                        slots[*dst] = Some(bn.forward(inp)?);
1352                    }
1353                }
1354
1355                Instruction::MultiHeadAttention {
1356                    input,
1357                    dst,
1358                    n_heads,
1359                } => {
1360                    let inp = get_slot(&slots, *input)?;
1361                    let d_model = *inp
1362                        .dims()
1363                        .last()
1364                        .ok_or_else(|| shrew_core::Error::msg("MHA input has no dimensions"))?;
1365                    let mut mhas = self.mha_blocks.borrow_mut();
1366                    if !mhas.contains_key(dst) {
1367                        let mha = shrew_nn::MultiHeadAttention::<B>::new(
1368                            d_model,
1369                            *n_heads,
1370                            inp.dtype(),
1371                            inp.device(),
1372                        )?;
1373                        mhas.insert(*dst, mha);
1374                    }
1375                    slots[*dst] = Some(mhas.get(dst).unwrap().forward(inp)?);
1376                }
1377
1378                Instruction::TransformerBlock {
1379                    input,
1380                    dst,
1381                    n_heads,
1382                } => {
1383                    let inp = get_slot(&slots, *input)?;
1384                    let dims = inp.dims();
1385                    let d_model = dims[dims.len() - 1];
1386                    let d_ff = d_model * 4;
1387                    let mut blocks = self.transformer_blocks.borrow_mut();
1388                    if !blocks.contains_key(dst) {
1389                        let block = shrew_nn::TransformerBlock::<B>::new(
1390                            d_model,
1391                            *n_heads,
1392                            d_ff,
1393                            true,
1394                            inp.dtype(),
1395                            inp.device(),
1396                        )?;
1397                        blocks.insert(*dst, block);
1398                    }
1399                    let block = blocks.get(dst).unwrap();
1400                    block.set_training(self.config.training);
1401                    slots[*dst] = Some(block.forward(inp)?);
1402                }
1403
1404
1405                Instruction::Dropout { src, dst, p } => {
1406                    let t = get_slot(&slots, *src)?;
1407                    let dropout = shrew_nn::Dropout::new(*p);
1408                    if self.config.training {
1409                        slots[*dst] = Some(dropout.forward_t(t)?);
1410                    } else {
1411                        slots[*dst] = Some(t.clone());
1412                    }
1413                }
1414
1415                Instruction::CrossEntropy {
1416                    predictions,
1417                    targets,
1418                    dst,
1419                } => {
1420                    let p = get_slot(&slots, *predictions)?;
1421                    let t = get_slot(&slots, *targets)?;
1422                    slots[*dst] = Some(shrew_nn::cross_entropy_loss(p, t)?);
1423                }
1424
1425                Instruction::MseLoss {
1426                    predictions,
1427                    targets,
1428                    dst,
1429                } => {
1430                    let p = get_slot(&slots, *predictions)?;
1431                    let t = get_slot(&slots, *targets)?;
1432                    slots[*dst] = Some(shrew_nn::mse_loss(p, t)?);
1433                }
1434
1435                Instruction::Constant {
1436                    value,
1437                    output_type,
1438                    dst,
1439                } => {
1440                    let tensor = materialize_constant::<B>(
1441                        value,
1442                        output_type,
1443                        self.config.default_dtype,
1444                        &self.device,
1445                    )?;
1446                    slots[*dst] = Some(tensor);
1447                }
1448
1449                Instruction::Repeat {
1450                    count,
1451                    body_op,
1452                    src,
1453                    dst,
1454                } => {
1455                    let t = get_slot(&slots, *src)?;
1456                    let mut current = t.clone();
1457                    for i in 0..(*count as usize) {
1458                        match body_op.as_ref() {
1459                            OpKind::TransformerBlock { n_heads } => {
1460                                let dims = current.dims();
1461                                let d_model = dims[dims.len() - 1];
1462                                let d_ff = d_model * 4;
1463                                let mut blocks = self.repeat_blocks.borrow_mut();
1464                                let key = (*dst, i);
1465                                if !blocks.contains_key(&key) {
1466                                    let block = shrew_nn::TransformerBlock::<B>::new(
1467                                        d_model,
1468                                        *n_heads as usize,
1469                                        d_ff,
1470                                        true,
1471                                        current.dtype(),
1472                                        current.device(),
1473                                    )?;
1474                                    blocks.insert(key, block);
1475                                }
1476                                let block = blocks.get(&key).unwrap();
1477                                block.set_training(self.config.training);
1478                                current = block.forward(&current)?;
1479                            }
1480                            _ => {
1481                                current = execute_body_op::<B>(body_op, &current, &self.device)?;
1482                            }
1483                        }
1484                    }
1485                    slots[*dst] = Some(current);
1486                }
1487
1488
1489                Instruction::Call {
1490                    graph_name,
1491                    inputs: input_slots,
1492                    dst,
1493                } => {
1494                    let _sub_cg = self.compiled.get(graph_name).ok_or_else(|| {
1495                        shrew_core::Error::msg(format!(
1496                            "Called graph '{}' not compiled",
1497                            graph_name
1498                        ))
1499                    })?;
1500                    let sub_graph = self.program.get_graph(graph_name).ok_or_else(|| {
1501                        shrew_core::Error::msg(format!("Called graph '{}' not found", graph_name))
1502                    })?;
1503                    let mut sub_inputs = HashMap::new();
1504                    for (i, &input_id) in sub_graph.inputs.iter().enumerate() {
1505                        if let Some(&s) = input_slots.get(i) {
1506                            if let Some(tensor) = &slots[s] {
1507                                let input_name = sub_graph.node(input_id).name.clone();
1508                                sub_inputs.insert(input_name, tensor.clone());
1509                            }
1510                        }
1511                    }
1512                    let result = self.run(graph_name, &sub_inputs)?;
1513                    if let Some(out) = result.output() {
1514                        slots[*dst] = Some(out.clone());
1515                    }
1516                }
1517
1518                Instruction::Compare { op, lhs, rhs, dst } => {
1519                    let a = get_slot(&slots, *lhs)?;
1520                    let b = get_slot(&slots, *rhs)?;
1521                    let result = match op {
1522                        CompareInstr::Equal => a.eq(b),
1523                        CompareInstr::NotEqual => a.ne(b),
1524                        CompareInstr::Less => a.lt(b),
1525                        CompareInstr::Greater => a.gt(b),
1526                        CompareInstr::LessEqual => a.le(b),
1527                        CompareInstr::GreaterEqual => a.ge(b),
1528                    }?;
1529                    slots[*dst] = Some(result);
1530                }
1531
1532                Instruction::LogicalNot { src, dst } => {
1533                    let t = get_slot(&slots, *src)?;
1534                    let data = t.to_f64_vec()?;
1535                    let result: Vec<f64> = data
1536                        .iter()
1537                        .map(|&v| if v == 0.0 { 1.0 } else { 0.0 })
1538                        .collect();
1539                    let n = result.len();
1540                    slots[*dst] = Some(Tensor::<B>::from_f64_slice(
1541                        &result,
1542                        n,
1543                        CoreDType::U8,
1544                        &self.device,
1545                    )?);
1546                }
1547
1548                Instruction::LogicalBinOp { op, lhs, rhs, dst } => {
1549                    let a = get_slot(&slots, *lhs)?;
1550                    let b = get_slot(&slots, *rhs)?;
1551                    let a_data = a.to_f64_vec()?;
1552                    let b_data = b.to_f64_vec()?;
1553                    let result: Vec<f64> = a_data
1554                        .iter()
1555                        .zip(b_data.iter())
1556                        .map(|(&x, &y)| match op {
1557                            LogicalBinInstr::And => {
1558                                if x != 0.0 && y != 0.0 {
1559                                    1.0
1560                                } else {
1561                                    0.0
1562                                }
1563                            }
1564                            LogicalBinInstr::Or => {
1565                                if x != 0.0 || y != 0.0 {
1566                                    1.0
1567                                } else {
1568                                    0.0
1569                                }
1570                            }
1571                        })
1572                        .collect();
1573                    let n = result.len();
1574                    slots[*dst] = Some(Tensor::<B>::from_f64_slice(
1575                        &result,
1576                        n,
1577                        CoreDType::U8,
1578                        &self.device,
1579                    )?);
1580                }
1581
1582                Instruction::FusedMatMulAdd { a, b, c, dst } => {
1583                    let at = get_slot(&slots, *a)?;
1584                    let bt = get_slot(&slots, *b)?;
1585                    let ct = get_slot(&slots, *c)?;
1586                    slots[*dst] = Some(at.matmul(bt)?.add(ct)?);
1587                }
1588
1589                Instruction::FusedAddRelu { lhs, rhs, dst } => {
1590                    let a = get_slot(&slots, *lhs)?;
1591                    let b = get_slot(&slots, *rhs)?;
1592                    slots[*dst] = Some(a.add(b)?.relu()?);
1593                }
1594
1595                Instruction::FusedSubRelu { lhs, rhs, dst } => {
1596                    let a = get_slot(&slots, *lhs)?;
1597                    let b = get_slot(&slots, *rhs)?;
1598                    slots[*dst] = Some(a.sub(b)?.relu()?);
1599                }
1600
1601                Instruction::FusedMatMulRelu { lhs, rhs, dst } => {
1602                    let a = get_slot(&slots, *lhs)?;
1603                    let b = get_slot(&slots, *rhs)?;
1604                    slots[*dst] = Some(a.matmul(b)?.relu()?);
1605                }
1606
1607                Instruction::Copy { src, dst } => {
1608                    let t = get_slot(&slots, *src)?;
1609                    slots[*dst] = Some(t.clone());
1610                }
1611
1612                Instruction::Range {
1613                    inputs: input_slots,
1614                    output_type,
1615                    dst,
1616                } => {
1617                    let (start, end) = if input_slots.len() >= 2 {
1618                        let s = get_slot(&slots, input_slots[0])?.to_scalar_f64()?;
1619                        let e = get_slot(&slots, input_slots[1])?.to_scalar_f64()?;
1620                        (s as i64, e as i64)
1621                    } else if input_slots.len() == 1 {
1622                        (
1623                            0i64,
1624                            get_slot(&slots, input_slots[0])?.to_scalar_f64()? as i64,
1625                        )
1626                    } else {
1627                        match output_type {
1628                            IrType::Tensor { shape, .. } => {
1629                                if let Some(Dim::Fixed(n)) = shape.first() {
1630                                    (0, *n)
1631                                } else {
1632                                    (0, 1)
1633                                }
1634                            }
1635                            _ => (0, 1),
1636                        }
1637                    };
1638                    let data: Vec<f64> = (start..end).map(|i| i as f64).collect();
1639                    let len = data.len();
1640                    slots[*dst] = Some(Tensor::<B>::from_f64_slice(
1641                        &data,
1642                        len,
1643                        CoreDType::I64,
1644                        &self.device,
1645                    )?);
1646                }
1647
1648                Instruction::Free { slot } => {
1649                    slots[*slot] = None;
1650                }
1651            }
1652        }
1653
1654        // Collect outputs
1655        let mut outputs = HashMap::new();
1656        for (name, &slot) in &cg.output_slots {
1657            if let Some(tensor) = &slots[slot] {
1658                outputs.insert(name.clone(), tensor.clone());
1659            }
1660        }
1661
1662        Ok(JitResult { outputs })
1663    }
1664
1665    /// Get the underlying program.
1666    pub fn program(&self) -> &IrProgram {
1667        &self.program
1668    }
1669
1670    /// Get runtime config.
1671    pub fn config(&self) -> &RuntimeConfig {
1672        &self.config
1673    }
1674
1675    /// Get all parameters.
1676    pub fn params(&self) -> &HashMap<(String, String), Tensor<B>> {
1677        &self.params
1678    }
1679
1680    /// Get parameters for a specific graph.
1681    pub fn graph_params(&self, graph_name: &str) -> Vec<Tensor<B>> {
1682        self.params
1683            .iter()
1684            .filter(|((g, _), _)| g == graph_name)
1685            .map(|(_, t)| t.clone())
1686            .collect()
1687    }
1688
1689    /// Update parameters after an optimizer step.
1690    pub fn update_params(&mut self, graph_name: &str, new_params: &[Tensor<B>]) {
1691        let param_names: Vec<String> = self
1692            .params
1693            .keys()
1694            .filter(|(g, _)| g == graph_name)
1695            .map(|(_, n)| n.clone())
1696            .collect();
1697
1698        for (name, tensor) in param_names.into_iter().zip(new_params.iter()) {
1699            self.params
1700                .insert((graph_name.to_string(), name), tensor.clone());
1701        }
1702    }
1703
1704    /// Recompile a single graph (e.g., after optimizer changes shapes).
1705    pub fn recompile(&mut self, graph_name: &str) -> Result<()> {
1706        let graph = self
1707            .program
1708            .get_graph(graph_name)
1709            .ok_or_else(|| shrew_core::Error::msg(format!("Graph '{}' not found", graph_name)))?;
1710        let cg = compile_graph(graph, &self.program, &self.config)?;
1711        self.compiled.insert(graph_name.to_string(), cg);
1712        Ok(())
1713    }
1714
1715    /// Dump the compiled instruction tape for a graph (debugging).
1716    pub fn dump(&self, graph_name: &str) -> Option<String> {
1717        let cg = self.compiled.get(graph_name)?;
1718        let mut out = format!("=== JIT Compiled: {} ===\n", cg.graph_name);
1719        out.push_str(&format!("{}\n\n", cg.stats));
1720        for (i, instr) in cg.instructions.iter().enumerate() {
1721            out.push_str(&format!("  [{:>3}] {:?}\n", i, instr));
1722        }
1723        out.push_str(&format!("\nOutputs: {:?}\n", cg.output_slots));
1724        Some(out)
1725    }
1726}
1727
1728// =============================================================================
1729// Helper functions
1730// =============================================================================
1731
1732/// Get a tensor from a buffer slot.
1733fn get_slot<B: Backend>(slots: &[Option<Tensor<B>>], idx: usize) -> Result<&Tensor<B>> {
1734    slots.get(idx).and_then(|s| s.as_ref()).ok_or_else(|| {
1735        shrew_core::Error::msg(format!(
1736            "Buffer slot {} is empty (value was freed or never produced)",
1737            idx
1738        ))
1739    })
1740}
1741
1742/// Resolve a negative dimension index.
1743fn resolve_neg_dim(dim: i64, rank: usize) -> usize {
1744    if dim < 0 {
1745        (rank as i64 + dim) as usize
1746    } else {
1747        dim as usize
1748    }
1749}
1750
1751/// Resolve a Vec<Dim> to concrete shape.
1752fn resolve_shape_vec(
1753    dims: &[Dim],
1754    config: &RuntimeConfig,
1755    program: &IrProgram,
1756) -> Result<Vec<usize>> {
1757    dims.iter()
1758        .map(|d| resolve_dim(d, config, program))
1759        .collect()
1760}
1761
1762/// Resolve a single Dim.
1763fn resolve_dim(dim: &Dim, config: &RuntimeConfig, program: &IrProgram) -> Result<usize> {
1764    match dim {
1765        Dim::Fixed(n) => Ok(*n as usize),
1766        Dim::Symbolic(name) => {
1767            if let Some(&val) = config.dims.get(name) {
1768                return Ok(val);
1769            }
1770            if let Some(shrew_ir::graph::ConfigValue::Int(n)) = program.config.get(name) {
1771                return Ok(*n as usize);
1772            }
1773            Err(shrew_core::Error::msg(format!(
1774                "Unresolved symbolic dimension: '{}'",
1775                name
1776            )))
1777        }
1778        Dim::Dynamic => Err(shrew_core::Error::msg(
1779            "Cannot resolve dynamic dimension at compile time",
1780        )),
1781    }
1782}
1783
1784/// Initialize a parameter tensor.
1785fn init_param<B: Backend>(
1786    ty: &IrType,
1787    init: &shrew_ir::graph::InitStrategy,
1788    frozen: bool,
1789    config: &RuntimeConfig,
1790    program: &IrProgram,
1791    device: &B::Device,
1792) -> Result<Tensor<B>> {
1793    let (shape, dtype) = resolve_type(ty, config, program)?;
1794
1795    let tensor = match init {
1796        shrew_ir::graph::InitStrategy::Zeros => Tensor::<B>::zeros(shape, dtype, device)?,
1797        shrew_ir::graph::InitStrategy::Ones => Tensor::<B>::ones(shape, dtype, device)?,
1798        shrew_ir::graph::InitStrategy::Normal { mean, std } => {
1799            Tensor::<B>::randn(shape, dtype, device)?.affine(*std, *mean)?
1800        }
1801        shrew_ir::graph::InitStrategy::Uniform { low, high } => {
1802            Tensor::<B>::rand(shape, dtype, device)?.affine(*high - *low, *low)?
1803        }
1804        shrew_ir::graph::InitStrategy::XavierUniform => {
1805            let (fan_in, fan_out) = compute_fans(&shape);
1806            let a = (6.0_f64 / (fan_in + fan_out) as f64).sqrt();
1807            Tensor::<B>::rand(shape, dtype, device)?.affine(2.0 * a, -a)?
1808        }
1809        shrew_ir::graph::InitStrategy::XavierNormal => {
1810            let (fan_in, fan_out) = compute_fans(&shape);
1811            let std = (2.0_f64 / (fan_in + fan_out) as f64).sqrt();
1812            Tensor::<B>::randn(shape, dtype, device)?.affine(std, 0.0)?
1813        }
1814        shrew_ir::graph::InitStrategy::KaimingUniform => {
1815            let (fan_in, _) = compute_fans(&shape);
1816            let bound = (3.0_f64 / fan_in as f64).sqrt();
1817            Tensor::<B>::rand(shape, dtype, device)?.affine(2.0 * bound, -bound)?
1818        }
1819        shrew_ir::graph::InitStrategy::KaimingNormal => {
1820            let (fan_in, _) = compute_fans(&shape);
1821            let std = (2.0_f64 / fan_in as f64).sqrt();
1822            Tensor::<B>::randn(shape, dtype, device)?.affine(std, 0.0)?
1823        }
1824        shrew_ir::graph::InitStrategy::Custom(_) => Tensor::<B>::randn(shape, dtype, device)?,
1825    };
1826
1827    if frozen {
1828        Ok(tensor)
1829    } else {
1830        Ok(tensor.set_variable())
1831    }
1832}
1833
1834/// Resolve IrType to (Shape, CoreDType).
1835fn resolve_type(
1836    ty: &IrType,
1837    config: &RuntimeConfig,
1838    program: &IrProgram,
1839) -> Result<(shrew_core::Shape, CoreDType)> {
1840    match ty {
1841        IrType::Tensor { shape, dtype } => {
1842            let dims: Vec<usize> = shape
1843                .iter()
1844                .map(|d| resolve_dim(d, config, program))
1845                .collect::<Result<Vec<_>>>()?;
1846            let core_dtype = ir_dtype_to_core(*dtype)?;
1847            Ok((shrew_core::Shape::new(dims), core_dtype))
1848        }
1849        IrType::Scalar(dtype) => {
1850            let core_dtype = ir_dtype_to_core(*dtype)?;
1851            Ok((shrew_core::Shape::new(vec![1]), core_dtype))
1852        }
1853        IrType::Int => Ok((shrew_core::Shape::new(vec![1]), CoreDType::I64)),
1854        _ => Ok((shrew_core::Shape::new(vec![1]), config.default_dtype)),
1855    }
1856}
1857
1858/// Materialize a constant as a tensor.
1859fn materialize_constant<B: Backend>(
1860    val: &ConstantValue,
1861    ty: &IrType,
1862    default_dtype: CoreDType,
1863    device: &B::Device,
1864) -> Result<Tensor<B>> {
1865    match val {
1866        ConstantValue::Int(n) => {
1867            Tensor::<B>::from_f64_slice(&[*n as f64], 1, CoreDType::I64, device)
1868        }
1869        ConstantValue::Float(f) => {
1870            let dtype = match ty {
1871                IrType::Tensor { dtype, .. } => ir_dtype_to_core(*dtype)?,
1872                IrType::Scalar(dtype) => ir_dtype_to_core(*dtype)?,
1873                _ => default_dtype,
1874            };
1875            Tensor::<B>::from_f64_slice(&[*f], 1, dtype, device)
1876        }
1877        ConstantValue::Bool(b) => {
1878            Tensor::<B>::from_f64_slice(&[if *b { 1.0 } else { 0.0 }], 1, CoreDType::U8, device)
1879        }
1880        ConstantValue::Str(_) => Tensor::<B>::zeros(1, default_dtype, device),
1881        ConstantValue::Null => Tensor::<B>::zeros(1, default_dtype, device),
1882    }
1883}
1884
1885/// Execute a body op (for Repeat instruction).
1886fn execute_body_op<B: Backend>(
1887    op: &OpKind,
1888    input: &Tensor<B>,
1889    _device: &B::Device,
1890) -> Result<Tensor<B>> {
1891    match op {
1892        OpKind::TransformerBlock { n_heads } => {
1893            let dims = input.dims();
1894            let d_model = dims[dims.len() - 1];
1895            let d_ff = d_model * 4;
1896            let block = shrew_nn::TransformerBlock::<B>::new(
1897                d_model,
1898                *n_heads as usize,
1899                d_ff,
1900                true,
1901                input.dtype(),
1902                input.device(),
1903            )?;
1904            block.forward(input)
1905        }
1906        OpKind::MultiHeadAttention { n_heads } => {
1907            let d_model = *input
1908                .dims()
1909                .last()
1910                .ok_or_else(|| shrew_core::Error::msg("MHA input has no dimensions"))?;
1911            let mha = shrew_nn::MultiHeadAttention::<B>::new(
1912                d_model,
1913                *n_heads as usize,
1914                input.dtype(),
1915                input.device(),
1916            )?;
1917            mha.forward(input)
1918        }
1919        _ => Err(shrew_core::Error::msg(format!(
1920            "Unsupported op in Repeat body: {:?}",
1921            op
1922        ))),
1923    }
1924}
1925
1926/// Compute (fan_in, fan_out) from a parameter shape.
1927fn compute_fans(shape: &shrew_core::Shape) -> (usize, usize) {
1928    let dims = shape.dims();
1929    match dims.len() {
1930        0 => (1, 1),
1931        1 => (dims[0], dims[0]),
1932        2 => (dims[1], dims[0]),
1933        _ => {
1934            let receptive: usize = dims[2..].iter().product();
1935            let fan_in = dims[1] * receptive;
1936            let fan_out = dims[0] * receptive;
1937            (fan_in, fan_out)
1938        }
1939    }
1940}
1941
1942// =============================================================================
1943// Convenience: parse → lower → validate → optimize → JIT compile
1944// =============================================================================
1945
1946/// Parse, lower, validate, optimize, and JIT-compile a `.sw` program.
1947///
1948/// This is the recommended entry point for production use — it produces
1949/// a JIT executor that runs graphs faster than the interpreter.
1950///
1951/// # Example
1952/// ```ignore
1953/// let jit = load_jit::<CpuBackend>(source, CpuDevice, RuntimeConfig::default())?;
1954/// let result = jit.run("Forward", &inputs)?;
1955/// ```
1956pub fn load_jit<B: Backend>(
1957    source: &str,
1958    device: B::Device,
1959    config: RuntimeConfig,
1960) -> Result<JitExecutor<B>> {
1961    let ast =
1962        shrew_ir::parse(source).map_err(|e| shrew_core::Error::msg(format!("Parse error: {e}")))?;
1963    let mut ir = shrew_ir::lower(&ast)
1964        .map_err(|e| shrew_core::Error::msg(format!("Lowering error: {e}")))?;
1965
1966    if let Err(errors) = shrew_ir::validate(&ir) {
1967        let msg = errors
1968            .iter()
1969            .map(|e| e.to_string())
1970            .collect::<Vec<_>>()
1971            .join("\n");
1972        return Err(shrew_core::Error::msg(format!("Validation errors:\n{msg}")));
1973    }
1974
1975    shrew_ir::infer_shapes(&mut ir);
1976    shrew_ir::optimize(&mut ir);
1977
1978    JitExecutor::<B>::compile(ir, device, config)
1979}