Skip to main content

shrew/exec/
engine.rs

1// =============================================================================
2// Engine — Core graph execution engine
3// =============================================================================
4//
5// Walks the IrGraph in topological order, dispatching each node to the
6// appropriate tensor operation. Manages parameter initialization and the
7// mapping from NodeId → live Tensor.
8
9use std::collections::HashMap;
10
11use shrew_core::backend::Backend;
12use shrew_core::dtype::DType as CoreDType;
13use shrew_core::error::Result;
14use shrew_core::tensor::Tensor;
15
16use shrew_ir::graph::{
17    ConfigValue, ConstantValue, DType as IrDType, Dim, InitStrategy, IrGraph, IrNode, IrProgram,
18    IrType, OpKind,
19};
20
21use shrew_nn::{
22    cross_entropy_loss, mse_loss, Dropout, Embedding, LayerNorm, Linear, Module, TransformerBlock,
23};
24
25// ─────────────────────────────────────────────────────────────────────────────
26// Runtime configuration
27// ─────────────────────────────────────────────────────────────────────────────
28
29/// Runtime configuration for resolving symbolic dimensions and execution mode.
30#[derive(Debug, Clone)]
31pub struct RuntimeConfig {
32    /// Maps symbolic dimension names to concrete values (e.g., "Batch" → 4).
33    pub dims: HashMap<String, usize>,
34    /// Default data type when unspecified (default: F32).
35    pub default_dtype: CoreDType,
36    /// Whether we're in training mode (affects dropout, etc.).
37    pub training: bool,
38}
39
40impl Default for RuntimeConfig {
41    fn default() -> Self {
42        Self {
43            dims: HashMap::new(),
44            default_dtype: CoreDType::F32,
45            training: false,
46        }
47    }
48}
49
50impl RuntimeConfig {
51    /// Set a symbolic dimension value.
52    pub fn set_dim(mut self, name: impl Into<String>, value: usize) -> Self {
53        self.dims.insert(name.into(), value);
54        self
55    }
56
57    /// Set training mode.
58    pub fn with_training(mut self, training: bool) -> Self {
59        self.training = training;
60        self
61    }
62
63    /// Set default dtype.
64    pub fn with_dtype(mut self, dtype: CoreDType) -> Self {
65        self.default_dtype = dtype;
66        self
67    }
68}
69
70// ─────────────────────────────────────────────────────────────────────────────
71// Execution result
72// ─────────────────────────────────────────────────────────────────────────────
73
74/// The result of executing a graph.
75#[derive(Debug)]
76pub struct ExecResult<B: Backend> {
77    /// Output tensors, keyed by node name.
78    pub outputs: HashMap<String, Tensor<B>>,
79    /// All intermediate values, keyed by NodeId.
80    pub values: HashMap<usize, Tensor<B>>,
81}
82
83impl<B: Backend> ExecResult<B> {
84    /// Get the first (or only) output tensor.
85    pub fn output(&self) -> Option<&Tensor<B>> {
86        self.outputs.values().next()
87    }
88
89    /// Get an output by node name.
90    pub fn get(&self, name: &str) -> Option<&Tensor<B>> {
91        self.outputs.get(name)
92    }
93}
94
95// ─────────────────────────────────────────────────────────────────────────────
96// Executor
97// ─────────────────────────────────────────────────────────────────────────────
98
99/// Executes IrProgram graphs on the Shrew tensor runtime.
100pub struct Executor<B: Backend> {
101    /// The lowered IR program.
102    program: IrProgram,
103    /// Runtime configuration (symbolic dims, dtype, training mode).
104    config: RuntimeConfig,
105    /// Device to execute on.
106    device: B::Device,
107    /// Initialized parameter tensors, keyed by (graph_name, param_name).
108    params: HashMap<(String, String), Tensor<B>>,
109    /// Persistent transformer blocks (e.g. from TransformerBlock or Repeat nodes).
110    transformer_blocks: std::sync::RwLock<HashMap<String, TransformerBlock<B>>>,
111    /// Persistent multi-head attention blocks.
112    mha_blocks: std::sync::RwLock<HashMap<String, shrew_nn::MultiHeadAttention<B>>>,
113}
114
115impl<B: Backend> Executor<B> {
116    /// Create a new executor. Initializes all parameters.
117    pub fn new(program: IrProgram, device: B::Device, config: RuntimeConfig) -> Result<Self> {
118        let mut exec = Self {
119            program,
120            config,
121            device,
122            params: HashMap::new(),
123            transformer_blocks: std::sync::RwLock::new(HashMap::new()),
124            mha_blocks: std::sync::RwLock::new(HashMap::new()),
125        };
126        exec.init_all_params()?;
127        Ok(exec)
128    }
129
130    /// Get the underlying IR program.
131    pub fn program(&self) -> &IrProgram {
132        &self.program
133    }
134
135    /// Get a reference to the runtime config.
136    pub fn config(&self) -> &RuntimeConfig {
137        &self.config
138    }
139
140    /// Get a mutable reference to the runtime config.
141    pub fn config_mut(&mut self) -> &mut RuntimeConfig {
142        &mut self.config
143    }
144
145    /// Get all parameter tensors.
146    pub fn params(&self) -> &HashMap<(String, String), Tensor<B>> {
147        &self.params
148    }
149
150    /// Get flattened parameter list (all graphs and persistent blocks).
151    pub fn all_params(&self) -> Vec<Tensor<B>> {
152        let mut list: Vec<Tensor<B>> = self.params.values().cloned().collect();
153        for tb in self.transformer_blocks.read().unwrap().values() {
154            list.extend(tb.parameters());
155        }
156        for mha in self.mha_blocks.read().unwrap().values() {
157            list.extend(mha.parameters());
158        }
159        list
160    }
161
162    /// Get all parameters as `(key, tensor)` pairs, where key = `"graph/param"`.
163    pub fn named_params(&self) -> Vec<(String, Tensor<B>)> {
164        let mut pairs: Vec<(String, Tensor<B>)> = self
165            .params
166            .iter()
167            .map(|((g, p), t)| (format!("{g}/{p}"), t.clone()))
168            .collect();
169        for (name, tb) in self.transformer_blocks.read().unwrap().iter() {
170            for (pname, param) in tb.named_parameters() {
171                pairs.push((format!("{name}/{pname}"), param));
172            }
173        }
174        for (name, mha) in self.mha_blocks.read().unwrap().iter() {
175            for (pname, param) in mha.named_parameters() {
176                pairs.push((format!("{name}/{pname}"), param));
177            }
178        }
179        pairs.sort_by(|a, b| a.0.cmp(&b.0));
180        pairs
181    }
182
183    /// Set a parameter by its `"graph/param"` key.  Returns true if found.
184    pub fn set_param_by_key(&mut self, key: &str, tensor: Tensor<B>) -> bool {
185        if let Some(pos) = key.find('/') {
186            let graph = &key[..pos];
187            let param = &key[pos + 1..];
188            let k = (graph.to_string(), param.to_string());
189            if let std::collections::hash_map::Entry::Occupied(mut e) = self.params.entry(k) {
190                e.insert(tensor.set_variable());
191                return true;
192            }
193        }
194        false
195    }
196
197    /// The device this executor is running on.
198    pub fn device(&self) -> &B::Device {
199        &self.device
200    }
201
202    /// Execute a named graph with given inputs.
203    pub fn run(
204        &self,
205        graph_name: &str,
206        inputs: &HashMap<String, Tensor<B>>,
207    ) -> Result<ExecResult<B>> {
208        let graph = self.program.get_graph(graph_name).ok_or_else(|| {
209            shrew_core::Error::msg(format!("Graph '{}' not found in program", graph_name))
210        })?;
211        self.execute_graph(graph, inputs)
212    }
213
214    /// Execute a graph, returning output tensors and all intermediate values.
215    fn execute_graph(
216        &self,
217        graph: &IrGraph,
218        inputs: &HashMap<String, Tensor<B>>,
219    ) -> Result<ExecResult<B>> {
220        let order = graph.topo_order();
221        let mut values: HashMap<usize, Tensor<B>> = HashMap::new();
222
223        // Map input nodes to their provided tensors
224        for &input_id in &graph.inputs {
225            let node = graph.node(input_id);
226            if let Some(tensor) = inputs.get(&node.name) {
227                values.insert(input_id.0, tensor.clone());
228            }
229        }
230
231        // Map parameter nodes to their initialized tensors
232        for param in &graph.params {
233            let key = (graph.name.clone(), param.name.clone());
234            if let Some(tensor) = self.params.get(&key) {
235                values.insert(param.node_id.0, tensor.clone());
236            }
237        }
238
239        // Execute each node in topological order
240        for &node_id in &order {
241            if values.contains_key(&node_id.0) {
242                continue; // Already initialized (input or param)
243            }
244            let node = graph.node(node_id);
245            let result = self.execute_node(graph, node, &values)?;
246            values.insert(node_id.0, result);
247        }
248
249        // Collect outputs
250        let mut outputs = HashMap::new();
251        for output in &graph.outputs {
252            if let Some(tensor) = values.get(&output.node_id.0) {
253                outputs.insert(output.name.clone(), tensor.clone());
254            }
255        }
256
257        Ok(ExecResult { outputs, values })
258    }
259
260    /// Execute a single node given its inputs' current values.
261    fn execute_node(
262        &self,
263        _graph: &IrGraph,
264        node: &IrNode,
265        values: &HashMap<usize, Tensor<B>>,
266    ) -> Result<Tensor<B>> {
267        // Collect input tensors for this node
268        let input_tensors: Vec<&Tensor<B>> = node
269            .inputs
270            .iter()
271            .filter_map(|id| values.get(&id.0))
272            .collect();
273
274        match &node.op {
275            // ── Identity: pass-through ──
276            OpKind::Identity => input_tensors.first().map(|t| (*t).clone()).ok_or_else(|| {
277                shrew_core::Error::msg(format!("Identity node '{}' has no input", node.name))
278            }),
279
280            // ── Unary ops ──
281            OpKind::Neg => unary(&input_tensors, &node.name, |t| t.neg()),
282            OpKind::Relu => unary(&input_tensors, &node.name, |t| t.relu()),
283            OpKind::Gelu => unary(&input_tensors, &node.name, |t| t.gelu()),
284            OpKind::Silu => unary(&input_tensors, &node.name, |t| t.silu()),
285            OpKind::Sigmoid => unary(&input_tensors, &node.name, |t| t.sigmoid()),
286            OpKind::Tanh => unary(&input_tensors, &node.name, |t| t.tanh()),
287            OpKind::Exp => unary(&input_tensors, &node.name, |t| t.exp()),
288            OpKind::Log => unary(&input_tensors, &node.name, |t| t.log()),
289            OpKind::Sqrt => unary(&input_tensors, &node.name, |t| t.sqrt()),
290
291            // ── Transpose ──
292            OpKind::Transpose => {
293                let t = require_input(&input_tensors, 0, &node.name)?;
294                let rank = t.rank();
295                if rank < 2 {
296                    return Err(shrew_core::Error::msg(format!(
297                        "Transpose requires rank >= 2, got {} for '{}'",
298                        rank, node.name
299                    )));
300                }
301                t.transpose(rank - 2, rank - 1)
302            }
303
304            // ── Binary ops ──
305            OpKind::Add => binary(&input_tensors, &node.name, |a, b| a.add(b)),
306            OpKind::Sub => binary(&input_tensors, &node.name, |a, b| a.sub(b)),
307            OpKind::Mul => binary(&input_tensors, &node.name, |a, b| a.mul(b)),
308            OpKind::Div => binary(&input_tensors, &node.name, |a, b| a.div(b)),
309            OpKind::MatMul => binary(&input_tensors, &node.name, |a, b| a.matmul(b)),
310
311            // ── Pow: x^y via exp(y * ln(x)) ──
312            OpKind::Pow => {
313                let base = require_input(&input_tensors, 0, &node.name)?;
314                let exp_t = require_input(&input_tensors, 1, &node.name)?;
315                // x^y = exp(y * ln(x))
316                base.log()?.mul(exp_t)?.exp()
317            }
318
319            // ── Mod: a - floor(a / b) * b ──
320            OpKind::Mod => {
321                let a = require_input(&input_tensors, 0, &node.name)?;
322                let b = require_input(&input_tensors, 1, &node.name)?;
323                let quotient = a.div(b)?.floor()?;
324                let product = quotient.mul(b)?;
325                a.sub(&product)
326            }
327
328            // ── Reduction ops ──
329            OpKind::Sum { dims, keepdim } => {
330                let t = require_input(&input_tensors, 0, &node.name)?;
331                if dims.is_empty() || (dims.len() == 1 && dims[0] == -1) {
332                    t.sum_all()
333                } else {
334                    let dim = resolve_neg_dim(dims[0], t.rank());
335                    t.sum(dim, *keepdim)
336                }
337            }
338
339            OpKind::Mean { dims, keepdim } => {
340                let t = require_input(&input_tensors, 0, &node.name)?;
341                if dims.is_empty() || (dims.len() == 1 && dims[0] == -1) {
342                    t.mean_all()
343                } else {
344                    let dim = resolve_neg_dim(dims[0], t.rank());
345                    t.mean(dim, *keepdim)
346                }
347            }
348
349            OpKind::Max { dim, keepdim } => {
350                let t = require_input(&input_tensors, 0, &node.name)?;
351                let d = resolve_neg_dim(*dim, t.rank());
352                t.max(d, *keepdim)
353            }
354
355            OpKind::Min { dim, keepdim } => {
356                let t = require_input(&input_tensors, 0, &node.name)?;
357                let d = resolve_neg_dim(*dim, t.rank());
358                t.min(d, *keepdim)
359            }
360
361            OpKind::Variance { dims, keepdim } => {
362                let t = require_input(&input_tensors, 0, &node.name)?;
363                if dims.is_empty() {
364                    t.var(0, *keepdim)
365                } else {
366                    let dim = resolve_neg_dim(dims[0], t.rank());
367                    t.var(dim, *keepdim)
368                }
369            }
370
371            // ── Softmax ──
372            OpKind::Softmax { dim } => {
373                let t = require_input(&input_tensors, 0, &node.name)?;
374                let d = resolve_neg_dim(*dim, t.rank());
375                t.softmax(d)
376            }
377
378            // ── Shape ops ──
379            OpKind::Reshape { target_shape } | OpKind::View { target_shape } => {
380                let t = require_input(&input_tensors, 0, &node.name)?;
381                let shape = self.resolve_shape_vec(target_shape)?;
382                t.reshape(shape)
383            }
384
385            OpKind::Permute { dims: perm_dims } => {
386                let t = require_input(&input_tensors, 0, &node.name)?;
387                // Apply successive transpositions to achieve the permutation
388                let mut result = t.clone();
389                let mut current: Vec<usize> = (0..t.rank()).collect();
390                for i in 0..perm_dims.len() {
391                    let target = perm_dims[i] as usize;
392                    if current[i] != target {
393                        let j = current.iter().position(|&x| x == target).ok_or_else(|| {
394                            shrew_core::Error::msg(format!(
395                                "permute: dimension {} not found in current layout",
396                                target
397                            ))
398                        })?;
399                        result = result.transpose(i, j)?;
400                        current.swap(i, j);
401                    }
402                }
403                Ok(result)
404            }
405
406            OpKind::Expand { target_shape } => {
407                let t = require_input(&input_tensors, 0, &node.name)?;
408                let shape = self.resolve_shape_vec(target_shape)?;
409                t.expand(shape)
410            }
411
412            OpKind::Concat { dim } => {
413                if input_tensors.is_empty() {
414                    return Err(shrew_core::Error::msg(format!(
415                        "Concat node '{}' has no inputs",
416                        node.name
417                    )));
418                }
419                let owned: Vec<Tensor<B>> = input_tensors.iter().map(|t| (*t).clone()).collect();
420                Tensor::<B>::cat(&owned, *dim as usize)
421            }
422
423            // ── Embedding ──
424            // Convention: embedding(indices, weight_table)
425            OpKind::Embedding => {
426                let indices = require_input(&input_tensors, 0, &node.name)?;
427                let table = require_input(&input_tensors, 1, &node.name)?;
428                let emb = Embedding::<B>::from_tensor(table.clone())?;
429                emb.forward(indices)
430            }
431
432            // ── Linear ──
433            // Convention: linear(input, weight) or linear(input, weight, bias)
434            OpKind::Linear { bias } => {
435                let input = require_input(&input_tensors, 0, &node.name)?;
436                let weight = require_input(&input_tensors, 1, &node.name)?;
437                if *bias && input_tensors.len() >= 3 {
438                    let bias_t = require_input(&input_tensors, 2, &node.name)?;
439                    let lin = Linear::<B>::from_tensors(weight.clone(), Some(bias_t.clone()))?;
440                    lin.forward(input)
441                } else {
442                    let lin = Linear::<B>::from_tensors(weight.clone(), None)?;
443                    lin.forward(input)
444                }
445            }
446
447            // ── LayerNorm ──
448            // Convention: layer_norm(input, weight, bias)
449            OpKind::LayerNorm { eps } => {
450                let input = require_input(&input_tensors, 0, &node.name)?;
451                let weight = require_input(&input_tensors, 1, &node.name)?;
452                let bias_t = require_input(&input_tensors, 2, &node.name)?;
453                let ln = LayerNorm::<B>::from_tensors(weight.clone(), bias_t.clone(), *eps)?;
454                ln.forward(input)
455            }
456
457            // ── MultiHeadAttention ──
458            OpKind::MultiHeadAttention { n_heads } => {
459                let input = require_input(&input_tensors, 0, &node.name)?;
460                let d_model = *input
461                    .dims()
462                    .last()
463                    .ok_or_else(|| shrew_core::Error::msg("MHA input has no dimensions"))?;
464                let mut mhas = self.mha_blocks.write().unwrap();
465                if !mhas.contains_key(&node.name) {
466                    let mha = shrew_nn::MultiHeadAttention::<B>::new(
467                        d_model,
468                        *n_heads as usize,
469                        input.dtype(),
470                        input.device(),
471                    )?;
472                    mhas.insert(node.name.clone(), mha);
473                }
474                mhas.get(&node.name).unwrap().forward(input)
475            }
476
477            // ── TransformerBlock ──
478            OpKind::TransformerBlock { n_heads } => {
479                let input = require_input(&input_tensors, 0, &node.name)?;
480                let dims = input.dims();
481                if dims.len() != 3 {
482                    return Err(shrew_core::Error::msg(format!(
483                        "TransformerBlock expects [batch, seq, d_model], got {:?}",
484                        dims
485                    )));
486                }
487                let d_model = dims[2];
488                let d_ff = d_model * 4;
489                let mut blocks = self.transformer_blocks.write().unwrap();
490                if !blocks.contains_key(&node.name) {
491                    let block = TransformerBlock::<B>::new(
492                        d_model,
493                        *n_heads as usize,
494                        d_ff,
495                        true, // causal by default
496                        input.dtype(),
497                        input.device(),
498                    )?;
499                    blocks.insert(node.name.clone(), block);
500                }
501                let block = blocks.get(&node.name).unwrap();
502                block.set_training(self.config.training);
503                block.forward(input)
504            }
505
506
507
508            // ── Dropout ──
509            OpKind::Dropout { p } => {
510                let input = require_input(&input_tensors, 0, &node.name)?;
511                let dropout = Dropout::new(*p);
512                if self.config.training {
513                    dropout.forward_t(input)
514                } else {
515                    Ok(input.clone())
516                }
517            }
518
519            // ── Loss functions ──
520            OpKind::CrossEntropy => {
521                let predictions = require_input(&input_tensors, 0, &node.name)?;
522                let targets = require_input(&input_tensors, 1, &node.name)?;
523                cross_entropy_loss(predictions, targets)
524            }
525
526            OpKind::MseLoss => {
527                let predictions = require_input(&input_tensors, 0, &node.name)?;
528                let targets = require_input(&input_tensors, 1, &node.name)?;
529                mse_loss(predictions, targets)
530            }
531
532            // ── Comparison ops ──
533            OpKind::Equal
534            | OpKind::NotEqual
535            | OpKind::Less
536            | OpKind::Greater
537            | OpKind::LessEqual
538            | OpKind::GreaterEqual => {
539                let lhs = require_input(&input_tensors, 0, &node.name)?;
540                let rhs = require_input(&input_tensors, 1, &node.name)?;
541                match &node.op {
542                    OpKind::Equal => lhs.eq(rhs),
543                    OpKind::NotEqual => lhs.ne(rhs),
544                    OpKind::Less => lhs.lt(rhs),
545                    OpKind::Greater => lhs.gt(rhs),
546                    OpKind::LessEqual => lhs.le(rhs),
547                    OpKind::GreaterEqual => lhs.ge(rhs),
548                    _ => unreachable!(),
549                }
550            }
551
552            // ── Constants ──
553            OpKind::Constant(val) => self.materialize_constant(val, &node.output_type),
554
555            // ── Repeat: execute body_op N times in sequence ──
556            OpKind::Repeat { count, body_op } => {
557                let input = require_input(&input_tensors, 0, &node.name)?;
558                let mut current = input.clone();
559                for i in 0..*count {
560                    let key = format!("{}/iter_{}", node.name, i);
561                    current = self.execute_body_op(body_op, &current, &key)?;
562                }
563                Ok(current)
564            }
565
566            // ── Call: execute another graph ──
567            OpKind::Call { graph_name } => {
568                // Build inputs for the sub-graph
569                let sub_graph = self.program.get_graph(graph_name).ok_or_else(|| {
570                    shrew_core::Error::msg(format!("Called graph '{}' not found", graph_name))
571                })?;
572                let mut sub_inputs = HashMap::new();
573                for (i, &input_id) in sub_graph.inputs.iter().enumerate() {
574                    let input_node = sub_graph.node(input_id);
575                    if let Some(tensor) = input_tensors.get(i) {
576                        sub_inputs.insert(input_node.name.clone(), (*tensor).clone());
577                    }
578                }
579                let result = self.execute_graph(sub_graph, &sub_inputs)?;
580                result.output().cloned().ok_or_else(|| {
581                    shrew_core::Error::msg(format!(
582                        "Called graph '{}' produced no output",
583                        graph_name
584                    ))
585                })
586            }
587
588            // ── Range ──
589            OpKind::Range => {
590                // range(start, end) → 1D tensor [start, start+1, ..., end-1]
591                let (start, end) = if input_tensors.len() >= 2 {
592                    let s = input_tensors[0].to_scalar_f64()?;
593                    let e = input_tensors[1].to_scalar_f64()?;
594                    (s as i64, e as i64)
595                } else if input_tensors.len() == 1 {
596                    (0i64, input_tensors[0].to_scalar_f64()? as i64)
597                } else {
598                    // Try resolving from output type shape
599                    match &node.output_type {
600                        IrType::Tensor { shape, .. } => {
601                            if let Some(Dim::Fixed(n)) = shape.first() {
602                                (0, *n)
603                            } else if let Some(Dim::Symbolic(name)) = shape.first() {
604                                let n = self.resolve_symbolic(name)? as i64;
605                                (0, n)
606                            } else {
607                                (0, 1)
608                            }
609                        }
610                        _ => (0, 1),
611                    }
612                };
613                let data: Vec<f64> = (start..end).map(|i| i as f64).collect();
614                let len = data.len();
615                Tensor::<B>::from_f64_slice(&data, len, CoreDType::I64, &self.device)
616            }
617
618            // ── BatchNorm ──
619            // Convention: batch_norm(input, weight, bias)
620            OpKind::BatchNorm { eps } => {
621                let input = require_input(&input_tensors, 0, &node.name)?;
622                if input_tensors.len() >= 3 {
623                    let weight = require_input(&input_tensors, 1, &node.name)?;
624                    let bias_t = require_input(&input_tensors, 2, &node.name)?;
625                    let bn = shrew_nn::BatchNorm2d::<B>::from_tensors(
626                        weight.clone(),
627                        bias_t.clone(),
628                        *eps,
629                    )?;
630                    bn.forward(input)
631                } else {
632                    // No weight/bias provided — create default BatchNorm from channels
633                    let dims = input.dims();
634                    if dims.len() != 4 {
635                        return Err(shrew_core::Error::msg(format!(
636                            "BatchNorm expects 4D input [N,C,H,W], got {:?}",
637                            dims
638                        )));
639                    }
640                    let c = dims[1];
641                    let bn =
642                        shrew_nn::BatchNorm2d::<B>::new(c, *eps, 0.1, input.dtype(), &self.device)?;
643                    bn.forward(input)
644                }
645            }
646
647            // ── Split ──
648            OpKind::Split { dim, chunks } => {
649                let input = require_input(&input_tensors, 0, &node.name)?;
650                let d = resolve_neg_dim(*dim, input.rank());
651                let result = input.chunk(*chunks as usize, d)?;
652                // Return first chunk (Split in IR produces a single node)
653                result
654                    .into_iter()
655                    .next()
656                    .ok_or_else(|| shrew_core::Error::msg("Split produced no chunks"))
657            }
658
659            // ── Logical ops (on comparison results) ──
660            OpKind::And => {
661                let lhs = require_input(&input_tensors, 0, &node.name)?;
662                let rhs = require_input(&input_tensors, 1, &node.name)?;
663                // a AND b = (a != 0) & (b != 0) → element-wise min
664                let a_data = lhs.to_f64_vec()?;
665                let b_data = rhs.to_f64_vec()?;
666                let result: Vec<f64> = a_data
667                    .iter()
668                    .zip(b_data.iter())
669                    .map(|(&a, &b)| if a != 0.0 && b != 0.0 { 1.0 } else { 0.0 })
670                    .collect();
671                let n = result.len();
672                Tensor::<B>::from_f64_slice(&result, n, CoreDType::U8, &self.device)
673            }
674            OpKind::Or => {
675                let lhs = require_input(&input_tensors, 0, &node.name)?;
676                let rhs = require_input(&input_tensors, 1, &node.name)?;
677                let a_data = lhs.to_f64_vec()?;
678                let b_data = rhs.to_f64_vec()?;
679                let result: Vec<f64> = a_data
680                    .iter()
681                    .zip(b_data.iter())
682                    .map(|(&a, &b)| if a != 0.0 || b != 0.0 { 1.0 } else { 0.0 })
683                    .collect();
684                let n = result.len();
685                Tensor::<B>::from_f64_slice(&result, n, CoreDType::U8, &self.device)
686            }
687            OpKind::Not => {
688                let input = require_input(&input_tensors, 0, &node.name)?;
689                let data = input.to_f64_vec()?;
690                let result: Vec<f64> = data
691                    .iter()
692                    .map(|&v| if v == 0.0 { 1.0 } else { 0.0 })
693                    .collect();
694                let n = result.len();
695                Tensor::<B>::from_f64_slice(&result, n, CoreDType::U8, &self.device)
696            }
697
698            // ── Custom op ──
699            OpKind::Custom { name, .. } => {
700                match name.as_str() {
701                    // Fused matmul + add: a.matmul(b) + c (no weight transpose)
702                    "fused_matmul_add" => {
703                        let a = require_input(&input_tensors, 0, &node.name)?;
704                        let b = require_input(&input_tensors, 1, &node.name)?;
705                        let c = require_input(&input_tensors, 2, &node.name)?;
706                        a.matmul(b)?.add(c)
707                    }
708                    // Fused add + relu
709                    "fused_add_relu" => {
710                        let a = require_input(&input_tensors, 0, &node.name)?;
711                        let b = require_input(&input_tensors, 1, &node.name)?;
712                        a.add(b)?.relu()
713                    }
714                    // Fused sub + relu
715                    "fused_sub_relu" => {
716                        let a = require_input(&input_tensors, 0, &node.name)?;
717                        let b = require_input(&input_tensors, 1, &node.name)?;
718                        a.sub(b)?.relu()
719                    }
720                    // Fused matmul + relu
721                    "fused_matmul_relu" => {
722                        let a = require_input(&input_tensors, 0, &node.name)?;
723                        let b = require_input(&input_tensors, 1, &node.name)?;
724                        a.matmul(b)?.relu()
725                    }
726                    _ => Err(shrew_core::Error::msg(format!(
727                        "Custom op '{}' is not implemented in the executor",
728                        name
729                    ))),
730                }
731            }
732        }
733    }
734
735    /// Execute a body op (used inside Repeat).
736    fn execute_body_op(&self, op: &OpKind, input: &Tensor<B>, key: &str) -> Result<Tensor<B>> {
737        match op {
738            OpKind::TransformerBlock { n_heads } => {
739                let dims = input.dims();
740                if dims.len() != 3 {
741                    return Err(shrew_core::Error::msg(format!(
742                        "TransformerBlock expects [batch, seq, d_model], got {:?}",
743                        dims
744                    )));
745                }
746                let d_model = dims[2];
747                let d_ff = d_model * 4;
748                let mut blocks = self.transformer_blocks.write().unwrap();
749                if !blocks.contains_key(key) {
750                    let block = TransformerBlock::<B>::new(
751                        d_model,
752                        *n_heads as usize,
753                        d_ff,
754                        true,
755                        input.dtype(),
756                        input.device(),
757                    )?;
758                    blocks.insert(key.to_string(), block);
759                }
760                let block = blocks.get(key).unwrap();
761                block.set_training(self.config.training);
762                block.forward(input)
763            }
764            OpKind::MultiHeadAttention { n_heads } => {
765                let d_model = *input
766                    .dims()
767                    .last()
768                    .ok_or_else(|| shrew_core::Error::msg("MHA input has no dimensions"))?;
769                let mut mhas = self.mha_blocks.write().unwrap();
770                if !mhas.contains_key(key) {
771                    let mha = shrew_nn::MultiHeadAttention::<B>::new(
772                        d_model,
773                        *n_heads as usize,
774                        input.dtype(),
775                        input.device(),
776                    )?;
777                    mhas.insert(key.to_string(), mha);
778                }
779                mhas.get(key).unwrap().forward(input)
780            }
781            // For other repeated ops, dispatch through the main execute_node
782            // infrastructure by returning an error so the caller knows.
783            _ => Err(shrew_core::Error::msg(format!(
784                "Unsupported op in Repeat body: {:?}. \
785                 Only TransformerBlock and MultiHeadAttention are supported.",
786                op
787            ))),
788        }
789    }
790
791
792    // ─────────────────────────────────────────────────────────────────────
793    // Parameter initialization
794    // ─────────────────────────────────────────────────────────────────────
795
796    /// Initialize all parameters across all graphs.
797    fn init_all_params(&mut self) -> Result<()> {
798        let graphs: Vec<(String, Vec<_>)> = self
799            .program
800            .graphs
801            .iter()
802            .map(|g| {
803                (
804                    g.name.clone(),
805                    g.params
806                        .iter()
807                        .map(|p| (p.name.clone(), p.ty.clone(), p.init.clone(), p.frozen))
808                        .collect::<Vec<_>>(),
809                )
810            })
811            .collect();
812
813        for (graph_name, params) in &graphs {
814            for (param_name, ty, init, frozen) in params {
815                let tensor = self.init_param(ty, init, *frozen)?;
816                self.params
817                    .insert((graph_name.clone(), param_name.clone()), tensor);
818            }
819        }
820        Ok(())
821    }
822
823    /// Initialize a single parameter tensor based on its type and init strategy.
824    fn init_param(&self, ty: &IrType, init: &InitStrategy, frozen: bool) -> Result<Tensor<B>> {
825        let (shape, dtype) = self.resolve_type(ty)?;
826        let tensor = match init {
827            InitStrategy::Zeros => Tensor::<B>::zeros(shape, dtype, &self.device)?,
828            InitStrategy::Ones => Tensor::<B>::ones(shape, dtype, &self.device)?,
829            InitStrategy::Normal { mean, std } => {
830                Tensor::<B>::randn(shape, dtype, &self.device)?.affine(*std, *mean)?
831            }
832            InitStrategy::Uniform { low, high } => {
833                let range = high - low;
834                Tensor::<B>::rand(shape, dtype, &self.device)?.affine(range, *low)?
835            }
836            InitStrategy::XavierUniform => {
837                // Xavier uniform: U(-a, a) where a = sqrt(6 / (fan_in + fan_out))
838                let (fan_in, fan_out) = compute_fans(&shape);
839                let a = (6.0_f64 / (fan_in + fan_out) as f64).sqrt();
840                Tensor::<B>::rand(shape, dtype, &self.device)?.affine(2.0 * a, -a)?
841            }
842            InitStrategy::XavierNormal => {
843                // Xavier normal: N(0, std) where std = sqrt(2 / (fan_in + fan_out))
844                let (fan_in, fan_out) = compute_fans(&shape);
845                let std = (2.0_f64 / (fan_in + fan_out) as f64).sqrt();
846                Tensor::<B>::randn(shape, dtype, &self.device)?.affine(std, 0.0)?
847            }
848            InitStrategy::KaimingUniform => {
849                // Kaiming uniform: U(-bound, bound) where bound = sqrt(3 / fan_in)
850                let (fan_in, _) = compute_fans(&shape);
851                let bound = (3.0_f64 / fan_in as f64).sqrt();
852                Tensor::<B>::rand(shape, dtype, &self.device)?.affine(2.0 * bound, -bound)?
853            }
854            InitStrategy::KaimingNormal => {
855                // Kaiming normal: N(0, std) where std = sqrt(2 / fan_in)
856                let (fan_in, _) = compute_fans(&shape);
857                let std = (2.0_f64 / fan_in as f64).sqrt();
858                Tensor::<B>::randn(shape, dtype, &self.device)?.affine(std, 0.0)?
859            }
860            InitStrategy::Custom(_) => Tensor::<B>::randn(shape, dtype, &self.device)?,
861        };
862
863        if frozen {
864            Ok(tensor)
865        } else {
866            Ok(tensor.set_variable())
867        }
868    }
869
870    /// Update parameters after an optimizer step.
871    pub fn update_params(&mut self, graph_name: &str, new_params: &[Tensor<B>]) {
872        let param_names: Vec<String> = self
873            .params
874            .keys()
875            .filter(|(g, _)| g == graph_name)
876            .map(|(_, n)| n.clone())
877            .collect();
878
879        for (name, tensor) in param_names.into_iter().zip(new_params.iter()) {
880            self.params
881                .insert((graph_name.to_string(), name), tensor.clone());
882        }
883    }
884
885    /// Collect parameters for a specific graph (for optimizer).
886    pub fn graph_params(&self, graph_name: &str) -> Vec<Tensor<B>> {
887        self.params
888            .iter()
889            .filter(|((g, _), _)| g == graph_name)
890            .map(|(_, t)| t.clone())
891            .collect()
892    }
893
894    // ─────────────────────────────────────────────────────────────────────
895    // Helpers
896    // ─────────────────────────────────────────────────────────────────────
897
898    /// Resolve a Dim to a concrete usize.
899    fn resolve_dim(&self, dim: &Dim) -> Result<usize> {
900        match dim {
901            Dim::Fixed(n) => Ok(*n as usize),
902            Dim::Symbolic(name) => self.resolve_symbolic(name),
903            Dim::Dynamic => Err(shrew_core::Error::msg(
904                "Cannot resolve dynamic dimension at runtime",
905            )),
906        }
907    }
908
909    /// Resolve a symbolic dimension name.
910    fn resolve_symbolic(&self, name: &str) -> Result<usize> {
911        // Try runtime config
912        if let Some(&val) = self.config.dims.get(name) {
913            return Ok(val);
914        }
915        // Try program config
916        if let Some(ConfigValue::Int(n)) = self.program.config.get(name) {
917            return Ok(*n as usize);
918        }
919        Err(shrew_core::Error::msg(format!(
920            "Unresolved symbolic dimension: '{}'. Set it via RuntimeConfig::set_dim()",
921            name
922        )))
923    }
924
925    /// Resolve an IrType to a concrete (Shape, CoreDType).
926    fn resolve_type(&self, ty: &IrType) -> Result<(shrew_core::Shape, CoreDType)> {
927        match ty {
928            IrType::Tensor { shape, dtype } => {
929                let dims: Vec<usize> = shape
930                    .iter()
931                    .map(|d| self.resolve_dim(d))
932                    .collect::<Result<Vec<_>>>()?;
933                let core_dtype = ir_dtype_to_core(*dtype)?;
934                Ok((shrew_core::Shape::new(dims), core_dtype))
935            }
936            IrType::Scalar(dtype) => {
937                let core_dtype = ir_dtype_to_core(*dtype)?;
938                Ok((shrew_core::Shape::new(vec![1]), core_dtype))
939            }
940            IrType::Int => Ok((shrew_core::Shape::new(vec![1]), CoreDType::I64)),
941            _ => Ok((shrew_core::Shape::new(vec![1]), self.config.default_dtype)),
942        }
943    }
944
945    /// Resolve a Vec<Dim> to a concrete shape tuple.
946    fn resolve_shape_vec(&self, dims: &[Dim]) -> Result<Vec<usize>> {
947        dims.iter().map(|d| self.resolve_dim(d)).collect()
948    }
949
950    /// Materialize a constant value as a tensor.
951    fn materialize_constant(&self, val: &ConstantValue, ty: &IrType) -> Result<Tensor<B>> {
952        match val {
953            ConstantValue::Int(n) => {
954                Tensor::<B>::from_f64_slice(&[*n as f64], 1, CoreDType::I64, &self.device)
955            }
956            ConstantValue::Float(f) => Tensor::<B>::from_f64_slice(
957                &[*f],
958                1,
959                ir_type_dtype(ty, self.config.default_dtype)?,
960                &self.device,
961            ),
962            ConstantValue::Bool(b) => Tensor::<B>::from_f64_slice(
963                &[if *b { 1.0 } else { 0.0 }],
964                1,
965                CoreDType::U8,
966                &self.device,
967            ),
968            ConstantValue::Str(_) => {
969                // Strings can't be tensors — return a dummy scalar
970                Tensor::<B>::zeros(1, self.config.default_dtype, &self.device)
971            }
972            ConstantValue::Null => Tensor::<B>::zeros(1, self.config.default_dtype, &self.device),
973        }
974    }
975}
976
977// ─────────────────────────────────────────────────────────────────────────────
978// Free helpers
979// ─────────────────────────────────────────────────────────────────────────────
980
981/// Convert IR DType to core DType.
982pub fn ir_dtype_to_core(dt: IrDType) -> Result<CoreDType> {
983    match dt {
984        IrDType::F32 => Ok(CoreDType::F32),
985        IrDType::F64 => Ok(CoreDType::F64),
986        IrDType::U8 => Ok(CoreDType::U8),
987        IrDType::U32 => Ok(CoreDType::U32),
988        IrDType::I64 => Ok(CoreDType::I64),
989        // Map unsupported types to closest supported
990        IrDType::F16 | IrDType::Bf16 => Ok(CoreDType::F32),
991        IrDType::I8 | IrDType::I16 | IrDType::I32 => Ok(CoreDType::I64),
992        IrDType::U16 => Ok(CoreDType::U32),
993        IrDType::U64 => Ok(CoreDType::U32),
994        IrDType::Bool => Ok(CoreDType::U8),
995        _ => Err(shrew_core::Error::msg(format!(
996            "Unsupported IR dtype: {dt}"
997        ))),
998    }
999}
1000
1001/// Extract dtype from IrType, with a fallback default.
1002fn ir_type_dtype(ty: &IrType, default: CoreDType) -> Result<CoreDType> {
1003    match ty {
1004        IrType::Tensor { dtype, .. } => ir_dtype_to_core(*dtype),
1005        IrType::Scalar(dtype) => ir_dtype_to_core(*dtype),
1006        _ => Ok(default),
1007    }
1008}
1009
1010/// Resolve a negative dimension index.
1011fn resolve_neg_dim(dim: i64, rank: usize) -> usize {
1012    if dim < 0 {
1013        (rank as i64 + dim) as usize
1014    } else {
1015        dim as usize
1016    }
1017}
1018
1019/// Require an input at a given index.
1020fn require_input<'a, B: Backend>(
1021    inputs: &[&'a Tensor<B>],
1022    idx: usize,
1023    node_name: &str,
1024) -> Result<&'a Tensor<B>> {
1025    inputs.get(idx).copied().ok_or_else(|| {
1026        shrew_core::Error::msg(format!(
1027            "Node '{}' expected input at index {}, but only {} inputs available",
1028            node_name,
1029            idx,
1030            inputs.len()
1031        ))
1032    })
1033}
1034
1035/// Execute a unary op.
1036fn unary<B: Backend>(
1037    inputs: &[&Tensor<B>],
1038    node_name: &str,
1039    f: impl FnOnce(&Tensor<B>) -> Result<Tensor<B>>,
1040) -> Result<Tensor<B>> {
1041    let t = require_input(inputs, 0, node_name)?;
1042    f(t)
1043}
1044
1045/// Execute a binary op.
1046fn binary<B: Backend>(
1047    inputs: &[&Tensor<B>],
1048    node_name: &str,
1049    f: impl FnOnce(&Tensor<B>, &Tensor<B>) -> Result<Tensor<B>>,
1050) -> Result<Tensor<B>> {
1051    let a = require_input(inputs, 0, node_name)?;
1052    let b = require_input(inputs, 1, node_name)?;
1053    f(a, b)
1054}
1055
1056/// Compute (fan_in, fan_out) from a parameter shape.
1057///
1058/// Follows PyTorch conventions:
1059/// - 1-D (bias): fan_in = fan_out = shape[0]
1060/// - 2-D (linear weight): fan_in = shape[1], fan_out = shape[0]
1061/// - 3-D+ (conv weight): fan_in = shape[1] * receptive, fan_out = shape[0] * receptive
1062fn compute_fans(shape: &shrew_core::Shape) -> (usize, usize) {
1063    let dims = shape.dims();
1064    match dims.len() {
1065        0 => (1, 1),
1066        1 => (dims[0], dims[0]),
1067        2 => (dims[1], dims[0]),
1068        _ => {
1069            // Conv: [out_channels, in_channels, *kernel_size]
1070            let receptive: usize = dims[2..].iter().product();
1071            let fan_in = dims[1] * receptive;
1072            let fan_out = dims[0] * receptive;
1073            (fan_in, fan_out)
1074        }
1075    }
1076}