1use 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#[derive(Debug, Clone)]
31pub struct RuntimeConfig {
32 pub dims: HashMap<String, usize>,
34 pub default_dtype: CoreDType,
36 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 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 pub fn with_training(mut self, training: bool) -> Self {
59 self.training = training;
60 self
61 }
62
63 pub fn with_dtype(mut self, dtype: CoreDType) -> Self {
65 self.default_dtype = dtype;
66 self
67 }
68}
69
70#[derive(Debug)]
76pub struct ExecResult<B: Backend> {
77 pub outputs: HashMap<String, Tensor<B>>,
79 pub values: HashMap<usize, Tensor<B>>,
81}
82
83impl<B: Backend> ExecResult<B> {
84 pub fn output(&self) -> Option<&Tensor<B>> {
86 self.outputs.values().next()
87 }
88
89 pub fn get(&self, name: &str) -> Option<&Tensor<B>> {
91 self.outputs.get(name)
92 }
93}
94
95pub struct Executor<B: Backend> {
101 program: IrProgram,
103 config: RuntimeConfig,
105 device: B::Device,
107 params: HashMap<(String, String), Tensor<B>>,
109 transformer_blocks: std::sync::RwLock<HashMap<String, TransformerBlock<B>>>,
111 mha_blocks: std::sync::RwLock<HashMap<String, shrew_nn::MultiHeadAttention<B>>>,
113}
114
115impl<B: Backend> Executor<B> {
116 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 pub fn program(&self) -> &IrProgram {
132 &self.program
133 }
134
135 pub fn config(&self) -> &RuntimeConfig {
137 &self.config
138 }
139
140 pub fn config_mut(&mut self) -> &mut RuntimeConfig {
142 &mut self.config
143 }
144
145 pub fn params(&self) -> &HashMap<(String, String), Tensor<B>> {
147 &self.params
148 }
149
150 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 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 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 pub fn device(&self) -> &B::Device {
199 &self.device
200 }
201
202 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 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 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 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 for &node_id in &order {
241 if values.contains_key(&node_id.0) {
242 continue; }
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 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 fn execute_node(
262 &self,
263 _graph: &IrGraph,
264 node: &IrNode,
265 values: &HashMap<usize, Tensor<B>>,
266 ) -> Result<Tensor<B>> {
267 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 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 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 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 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 OpKind::Pow => {
313 let base = require_input(&input_tensors, 0, &node.name)?;
314 let exp_t = require_input(&input_tensors, 1, &node.name)?;
315 base.log()?.mul(exp_t)?.exp()
317 }
318
319 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 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 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 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 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 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 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 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 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 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, 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 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 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 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 OpKind::Constant(val) => self.materialize_constant(val, &node.output_type),
554
555 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, ¤t, &key)?;
562 }
563 Ok(current)
564 }
565
566 OpKind::Call { graph_name } => {
568 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 OpKind::Range => {
590 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 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 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 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 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 result
654 .into_iter()
655 .next()
656 .ok_or_else(|| shrew_core::Error::msg("Split produced no chunks"))
657 }
658
659 OpKind::And => {
661 let lhs = require_input(&input_tensors, 0, &node.name)?;
662 let rhs = require_input(&input_tensors, 1, &node.name)?;
663 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 OpKind::Custom { name, .. } => {
700 match name.as_str() {
701 "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" => {
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" => {
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" => {
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 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 _ => 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 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 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 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 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 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 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 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 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 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 fn resolve_symbolic(&self, name: &str) -> Result<usize> {
911 if let Some(&val) = self.config.dims.get(name) {
913 return Ok(val);
914 }
915 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 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 fn resolve_shape_vec(&self, dims: &[Dim]) -> Result<Vec<usize>> {
947 dims.iter().map(|d| self.resolve_dim(d)).collect()
948 }
949
950 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 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
977pub 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 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
1001fn 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
1010fn 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
1019fn 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
1035fn 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
1045fn 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
1056fn 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 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}