1use 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#[derive(Debug, Clone)]
59pub struct SlotRef {
60 pub slot: usize,
62}
63
64#[derive(Debug, Clone)]
69pub enum Instruction {
70 LoadInput {
73 name: String,
75 dst: usize,
77 },
78 LoadParam {
80 graph_name: String,
82 param_name: String,
83 dst: usize,
85 },
86
87 Unary {
89 op: UnaryInstr,
90 src: usize,
91 dst: usize,
92 },
93
94 Binary {
96 op: BinaryInstr,
97 lhs: usize,
98 rhs: usize,
99 dst: usize,
100 },
101
102 Reduce {
104 op: ReduceInstr,
105 src: usize,
106 dst: usize,
107 dims: Vec<i64>,
108 keepdim: bool,
109 },
110
111 Reshape {
113 src: usize,
114 dst: usize,
115 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 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 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 Constant {
205 value: ConstantValue,
206 output_type: IrType,
207 dst: usize,
208 },
209
210 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 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 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 Copy {
266 src: usize,
267 dst: usize,
268 },
269
270 Range {
272 inputs: Vec<usize>,
273 output_type: IrType,
274 dst: usize,
275 },
276
277 Free {
279 slot: usize,
280 },
281}
282
283#[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#[derive(Debug, Clone, Copy)]
299pub enum BinaryInstr {
300 Add,
301 Sub,
302 Mul,
303 Div,
304 MatMul,
305 Pow,
306 Mod,
307}
308
309#[derive(Debug, Clone, Copy)]
311pub enum ReduceInstr {
312 Sum,
313 Mean,
314 Max,
315 Min,
316 Variance,
317}
318
319#[derive(Debug, Clone, Copy)]
321pub enum CompareInstr {
322 Equal,
323 NotEqual,
324 Less,
325 Greater,
326 LessEqual,
327 GreaterEqual,
328}
329
330#[derive(Debug, Clone, Copy)]
332pub enum LogicalBinInstr {
333 And,
334 Or,
335}
336
337#[derive(Debug, Clone)]
343pub struct ValueLifetime {
344 pub produced_at: usize,
346 pub last_used_at: usize,
348 pub node_id: NodeId,
350 pub is_output: bool,
352 pub is_external: bool,
354}
355
356#[derive(Debug, Clone)]
359pub struct MemoryPlan {
360 pub num_slots: usize,
362 pub node_to_slot: HashMap<usize, usize>,
364 pub lifetimes: Vec<ValueLifetime>,
366 pub free_points: Vec<(usize, usize)>,
368 pub reuse_count: usize,
370}
371
372#[derive(Debug)]
381pub struct CompiledGraph {
382 pub graph_name: String,
384 pub instructions: Vec<Instruction>,
386 pub memory_plan: MemoryPlan,
388 pub output_slots: HashMap<String, usize>,
390 pub num_slots: usize,
392 pub stats: CompileStats,
394}
395
396#[derive(Debug, Clone)]
398pub struct CompileStats {
399 pub num_instructions: usize,
401 pub num_source_nodes: usize,
403 pub num_slots: usize,
405 pub num_reused: usize,
407 pub num_frees: usize,
409 pub num_fused: usize,
411 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
431pub fn compile_graph(
437 graph: &IrGraph,
438 program: &IrProgram,
439 config: &RuntimeConfig,
440) -> Result<CompiledGraph> {
441 let start = Instant::now();
442
443 let order = graph.topo_order();
445
446 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 let output_node_ids: std::collections::HashSet<usize> =
454 graph.outputs.iter().map(|o| o.node_id.0).collect();
455
456 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 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 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 produced_at.insert(node_id.0, instr_idx);
479
480 for &input_id in &node.inputs {
482 last_used_at.insert(input_id.0, instr_idx);
483 }
484
485 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 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 let instr = compile_node(graph, node, &node_to_slot, config, program)?;
508
509 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 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 if !is_output && !is_external && last < instructions.len().saturating_sub(1) {
545 free_points.push((slot, last));
546 }
547 }
548
549 free_points.sort_by(|a, b| b.1.cmp(&a.1));
551
552 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 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(), reuse_count: 0, };
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
598fn 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 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 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 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 OpKind::Transpose => Ok(Instruction::Transpose { src: slot(0)?, dst }),
722
723 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 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 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), chunks: *chunks as usize,
807 }),
808
809 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 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 OpKind::Constant(val) => Ok(Instruction::Constant {
884 value: val.clone(),
885 output_type: node.output_type.clone(),
886 dst,
887 }),
888
889 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 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 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 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 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 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
1005pub struct JitExecutor<B: Backend> {
1024 compiled: HashMap<String, CompiledGraph>,
1026 program: IrProgram,
1028 config: RuntimeConfig,
1030 device: B::Device,
1032 params: HashMap<(String, String), Tensor<B>>,
1034 transformer_blocks: std::cell::RefCell<HashMap<usize, shrew_nn::TransformerBlock<B>>>,
1036 mha_blocks: std::cell::RefCell<HashMap<usize, shrew_nn::MultiHeadAttention<B>>>,
1038 repeat_blocks: std::cell::RefCell<HashMap<(usize, usize), shrew_nn::TransformerBlock<B>>>,
1040}
1041
1042#[derive(Debug)]
1044pub struct JitResult<B: Backend> {
1045 pub outputs: HashMap<String, Tensor<B>>,
1047}
1048
1049impl<B: Backend> JitResult<B> {
1050 pub fn output(&self) -> Option<&Tensor<B>> {
1052 self.outputs.values().next()
1053 }
1054
1055 pub fn get(&self, name: &str) -> Option<&Tensor<B>> {
1057 self.outputs.get(name)
1058 }
1059}
1060
1061impl<B: Backend> JitExecutor<B> {
1062 pub fn compile(program: IrProgram, device: B::Device, config: RuntimeConfig) -> Result<Self> {
1064 let mut compiled = HashMap::new();
1065
1066 for graph in &program.graphs {
1068 let cg = compile_graph(graph, &program, &config)?;
1069 compiled.insert(graph.name.clone(), cg);
1070 }
1071
1072 let mut params = HashMap::new();
1074 for graph in &program.graphs {
1075 for param in &graph.params {
1076 let tensor = init_param::<B>(
1077 ¶m.ty,
1078 ¶m.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 pub fn stats(&self, graph_name: &str) -> Option<&CompileStats> {
1103 self.compiled.get(graph_name).map(|cg| &cg.stats)
1104 }
1105
1106 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 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 let mut slots: Vec<Option<Tensor<B>>> = vec![None; cg.num_slots];
1130
1131 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(), 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(¤t)?;
1479 }
1480 _ => {
1481 current = execute_body_op::<B>(body_op, ¤t, &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 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 pub fn program(&self) -> &IrProgram {
1667 &self.program
1668 }
1669
1670 pub fn config(&self) -> &RuntimeConfig {
1672 &self.config
1673 }
1674
1675 pub fn params(&self) -> &HashMap<(String, String), Tensor<B>> {
1677 &self.params
1678 }
1679
1680 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 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 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 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
1728fn 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
1742fn 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
1751fn 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
1762fn 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
1784fn 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
1834fn 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
1858fn 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
1885fn 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
1926fn 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
1942pub 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}