1use std::collections::HashMap;
25
26use crate::ast::SlotShape;
27use crate::ast::{CompiledU64Op, PolydatNode};
28use crate::kernel::WireSource;
29
30#[cfg(feature = "jit")]
31use crate::compile::jit::{self, JitOp};
32
33enum HybridStep {
35 #[cfg(feature = "jit")]
38 Jit(JitSegment),
39 Closure(ClosureStep),
41}
42
43#[cfg(feature = "jit")]
44struct JitSegment {
45 code_fn: crate::compile::jit::NativeFn,
46 _module: crate::compile::jit::JitCode,
49 fallible: bool,
52 input_slots: Vec<usize>,
55 output_slots: Vec<usize>,
56 nodes: Vec<usize>,
59}
60
61#[cfg(feature = "jit")]
64struct PendingSegment {
65 step: usize,
66 batch: Vec<(JitOp, Vec<usize>, Vec<usize>)>,
67 input_slots: Vec<usize>,
68 output_slots: Vec<usize>,
69 nodes: Vec<usize>,
70}
71
72impl HybridStep {
73 fn input_slots(&self) -> &[usize] {
74 match self {
75 #[cfg(feature = "jit")]
76 HybridStep::Jit(seg) => &seg.input_slots,
77 HybridStep::Closure(cs) => &cs.input_slots,
78 }
79 }
80 fn output_slots(&self) -> &[usize] {
81 match self {
82 #[cfg(feature = "jit")]
83 HybridStep::Jit(seg) => &seg.output_slots,
84 HybridStep::Closure(cs) => &cs.output_slots,
85 }
86 }
87 fn accepts_none(&self) -> bool {
90 match self {
91 #[cfg(feature = "jit")]
92 HybridStep::Jit(_) => false,
93 HybridStep::Closure(cs) => cs.accepts_none,
94 }
95 }
96
97 #[cfg_attr(not(feature = "jit"), allow(unused_variables))]
100 fn failing_node(&self, buffer: &[u64], tracker: usize) -> usize {
101 match self {
102 #[cfg(feature = "jit")]
103 HybridStep::Jit(seg) => seg
104 .nodes
105 .get(buffer[tracker] as usize)
106 .copied()
107 .unwrap_or(usize::MAX),
108 HybridStep::Closure(cs) => cs.node,
109 }
110 }
111}
112
113enum ClosureOp {
117 U64(CompiledU64Op),
118 Slot(crate::ast::CompiledSlotOp),
119}
120
121struct ClosureStep {
122 op: ClosureOp,
123 input_slots: Vec<usize>,
124 output_slots: Vec<usize>,
125 scratch_range: (usize, usize),
127 accepts_none: bool,
129 node: usize,
131}
132
133type ResolvedOutput = (
136 usize,
137 crate::ast::PortType,
138 Option<std::sync::Arc<[usize]>>,
139 bool,
140);
141
142struct HybridCore {
148 engine: crate::compile::select::Engine,
153 buffer: Vec<u64>,
154 coord_count: usize,
155 steps: std::sync::Arc<Vec<HybridStep>>,
156 output_map: HashMap<String, usize>,
157 gather_buf: Vec<u64>,
158 scatter_buf: Vec<u64>,
159 scratch: Vec<crate::ast::ScratchBuf>,
163 ref_slots: Vec<bool>,
165 ref_scratch: Vec<(usize, usize)>,
167 output_types: HashMap<String, crate::ast::PortType>,
169 externs: crate::compile::externs::Externs,
171 traversals: std::sync::Arc<[crate::dsl::traversal::Traversal]>,
174 resolved_outputs: Vec<Option<ResolvedOutput>>,
177 _nodes: std::sync::Arc<Vec<Box<dyn PolydatNode>>>,
179 drive: crate::compile::Drive,
183 none: Vec<bool>,
185 ran: Vec<u64>,
188 epoch: u64,
194 all_ran: bool,
196 clean: Vec<bool>,
200 use_clean: bool,
202 plan: std::sync::Arc<crate::compile::Invalidation>,
205 volatile: std::sync::Arc<[bool]>,
207 side: std::sync::Arc<[bool]>,
209 slot_step: std::sync::Arc<[Option<usize>]>,
211 sites: std::sync::Arc<crate::compile::Attribution>,
213 cur_step: usize,
215 tracker: usize,
218 all: std::sync::Arc<[usize]>,
220 dirty: std::sync::Arc<[Vec<usize>]>,
226 any_none: bool,
230 volatile_steps: std::sync::Arc<[usize]>,
232}
233
234impl Clone for HybridCore {
235 fn clone(&self) -> Self {
236 let mut core = HybridCore {
237 engine: self.engine,
238 buffer: self.buffer.clone(),
239 coord_count: self.coord_count,
240 steps: self.steps.clone(),
241 output_map: self.output_map.clone(),
242 gather_buf: self.gather_buf.clone(),
243 scatter_buf: self.scatter_buf.clone(),
244 scratch: self.scratch.clone(),
245 ref_slots: self.ref_slots.clone(),
246 ref_scratch: self.ref_scratch.clone(),
247 output_types: self.output_types.clone(),
248 externs: self.externs.clone(),
249 traversals: self.traversals.clone(),
250 resolved_outputs: self.resolved_outputs.clone(),
251 _nodes: self._nodes.clone(),
252 drive: self.drive.clone(),
253 none: self.none.clone(),
254 ran: self.ran.clone(),
255 epoch: self.epoch,
256 all_ran: self.all_ran,
257 clean: self.clean.clone(),
258 use_clean: self.use_clean,
259 plan: self.plan.clone(),
260 volatile: self.volatile.clone(),
261 side: self.side.clone(),
262 slot_step: self.slot_step.clone(),
263 sites: self.sites.clone(),
264 cur_step: self.cur_step,
265 tracker: self.tracker,
266 all: self.all.clone(),
267 dirty: self.dirty.clone(),
268 any_none: self.any_none,
269 volatile_steps: self.volatile_steps.clone(),
270 };
271 core.republish_refs();
272 core
273 }
274}
275
276impl HybridCore {
277 crate::compile::shared_core_methods!();
278}
279
280impl HybridCore {
281 fn set_use_clean(&mut self, on: bool) {
285 self.use_clean = on;
286 let side = std::sync::Arc::clone(&self.side);
287 self.dirty = self
288 .plan
289 .input_dependents
290 .iter()
291 .map(|deps| {
292 if on {
293 deps.clone()
294 } else {
295 deps.iter().copied().filter(|&i| side[i]).collect()
296 }
297 })
298 .collect::<Vec<_>>()
299 .into();
300 }
301
302 #[inline]
306 fn step_can_fail(&self, i: usize) -> bool {
307 match &self.steps[i] {
308 #[cfg(feature = "jit")]
309 HybridStep::Jit(seg) => seg.fallible,
310 _ => true,
311 }
312 }
313
314 #[inline]
320 fn failing_node(&self) -> usize {
321 self.steps[self.cur_step].failing_node(&self.buffer, self.tracker)
322 }
323
324 #[inline]
326 fn run_order(&mut self, order: &[usize]) {
327 let steps = &self.steps;
328 let none_free = !self.any_none;
329 for &i in order {
330 if self.all_ran || self.ran[i] == self.epoch {
331 continue;
332 }
333 let never = self.volatile[i];
334 if (self.use_clean || self.side[i]) && self.clean[i] && !never {
335 self.ran[i] = self.epoch;
336 continue;
337 }
338 self.cur_step = i;
339 run_hybrid_step(
340 &steps[i],
341 none_free,
342 &mut self.buffer,
343 &mut self.none,
344 &mut self.gather_buf,
345 &mut self.scatter_buf,
346 &mut self.scratch,
347 );
348 self.ran[i] = self.epoch;
349 self.clean[i] = !never;
350 }
351 }
352
353 #[inline]
359 fn run_fresh(&mut self) {
360 let steps = &self.steps;
361 for (i, step) in steps.iter().enumerate() {
362 if self.side[i] {
363 let never = self.volatile[i];
364 if self.clean[i] && !never {
365 continue;
366 }
367 self.clean[i] = !never;
368 }
369 self.cur_step = i;
370 run_hybrid_step(
371 step,
372 true,
373 &mut self.buffer,
374 &mut self.none,
375 &mut self.gather_buf,
376 &mut self.scatter_buf,
377 &mut self.scratch,
378 );
379 }
380 self.all_ran = true;
381 }
382
383 fn plan(&self) -> crate::EnginePlan {
385 let (native_segments, closure_steps) = self.engine_counts();
386 crate::EnginePlan {
387 native_segments,
388 closure_steps,
389 interpreted_nodes: 0,
390 }
391 }
392}
393
394#[inline]
397fn eval_all_hybrid_steps(core: &mut HybridCore) {
398 core.drive.stale = true;
399 core.eval_all();
400}
401
402impl HybridCore {
407 fn engine_counts(&self) -> (usize, usize) {
410 let closures = self
411 .steps
412 .iter()
413 .filter(|s| matches!(s, HybridStep::Closure(_)))
414 .count();
415 (self.steps.len() - closures, closures)
416 }
417}
418
419#[derive(Clone)]
424pub struct HybridKernelRaw {
425 core: HybridCore,
426}
427
428impl HybridKernelRaw {
429 crate::compile::kernel_accessors!(set_coords);
430 #[inline]
433 fn set_coords(&mut self, coords: &[u64]) {
434 for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
435 if self.core.buffer[i] != c {
436 self.core.buffer[i] = c;
437 self.core.dirty_input(i);
438 }
439 }
440 }
441
442 #[inline]
444 pub fn eval(&mut self, coords: &[u64]) {
445 self.set_coords(coords);
446 eval_all_hybrid_steps(&mut self.core);
447 }
448
449 pub fn set_input(
453 &mut self,
454 name: &str,
455 value: crate::ast::Value,
456 ) -> Result<(), crate::kernel::WriteError> {
457 self.core.set_extern(name, value).map(|_| ())
458 }
459
460 pub fn set_input_at(
462 &mut self,
463 index: usize,
464 value: crate::ast::Value,
465 ) -> Result<(), crate::kernel::WriteError> {
466 self.core.set_extern_at(index, value).map(|_| ())
467 }
468
469 #[inline]
471 pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
472 self.core.guard_ref_slot(slot);
473 self.eval(coords);
474 self.core.buffer[slot]
475 }
476
477 pub fn engine_counts(&self) -> (usize, usize) {
480 self.core.engine_counts()
481 }
482
483 pub fn retain_nodes(&mut self, nodes: Vec<Box<dyn PolydatNode>>) {
485 self.core._nodes = std::sync::Arc::new(nodes);
486 }
487}
488
489#[derive(Clone)]
501pub struct HybridKernelPull {
502 core: HybridCore,
503 slot_provenance: Vec<crate::kernel::ProvMask>,
504 changed_mask: crate::kernel::ProvMask,
505 force_run: bool,
508}
509
510impl HybridKernelPull {
511 crate::compile::kernel_accessors!(set_inputs);
512 #[inline]
515 fn set_inputs(&mut self, coords: &[u64]) {
516 self.changed_mask.clear();
517 for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
518 if self.core.buffer[i] != c {
519 self.core.buffer[i] = c;
520 self.changed_mask.set(i);
521 self.core.dirty_input(i);
522 }
523 }
524 }
525
526 #[inline]
528 pub fn eval(&mut self, coords: &[u64]) {
529 self.set_inputs(coords);
530 self.force_run = false;
531 eval_all_hybrid_steps(&mut self.core);
532 }
533
534 #[inline]
537 pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
538 self.core.guard_ref_slot(slot);
539 self.set_inputs(coords);
540 if !self.force_run
541 && slot < self.slot_provenance.len()
542 && !self.slot_provenance[slot].intersects(&self.changed_mask)
543 {
544 return self.core.buffer[slot];
545 }
546 self.force_run = false;
547 eval_all_hybrid_steps(&mut self.core);
548 self.core.buffer[slot]
549 }
550
551 pub fn set_input(
555 &mut self,
556 name: &str,
557 value: crate::ast::Value,
558 ) -> Result<(), crate::kernel::WriteError> {
559 self.core.set_extern(name, value)?;
560 self.force_run = true;
561 Ok(())
562 }
563
564 pub fn set_input_at(
566 &mut self,
567 index: usize,
568 value: crate::ast::Value,
569 ) -> Result<(), crate::kernel::WriteError> {
570 self.core.set_extern_at(index, value)?;
571 self.force_run = true;
572 Ok(())
573 }
574
575 pub fn engine_counts(&self) -> (usize, usize) {
578 self.core.engine_counts()
579 }
580
581 pub fn retain_nodes(&mut self, nodes: Vec<Box<dyn PolydatNode>>) {
583 self.core._nodes = std::sync::Arc::new(nodes);
584 }
585}
586
587#[derive(Clone)]
601pub struct HybridKernelPushPull {
602 core: HybridCore,
603 slot_provenance: Vec<crate::kernel::ProvMask>,
604 changed_mask: crate::kernel::ProvMask,
605 force_run: bool,
608}
609
610impl HybridKernelPushPull {
611 crate::compile::kernel_accessors!(set_inputs);
612 pub fn set_input(
617 &mut self,
618 name: &str,
619 value: crate::ast::Value,
620 ) -> Result<(), crate::kernel::WriteError> {
621 self.core.set_extern(name, value)?;
622 self.force_run = true;
623 Ok(())
624 }
625
626 pub fn set_input_at(
628 &mut self,
629 index: usize,
630 value: crate::ast::Value,
631 ) -> Result<(), crate::kernel::WriteError> {
632 self.core.set_extern_at(index, value)?;
633 self.force_run = true;
634 Ok(())
635 }
636
637 #[inline]
639 fn set_inputs(&mut self, coords: &[u64]) {
640 self.changed_mask.clear();
641 for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
642 if self.core.buffer[i] != c {
643 self.core.buffer[i] = c;
644 self.changed_mask.set(i);
645 self.core.dirty_input(i);
646 }
647 }
648 }
649
650 #[inline]
652 pub fn eval(&mut self, coords: &[u64]) {
653 self.set_inputs(coords);
654 self.force_run = false;
655 self.core.drive.stale = true;
656 self.core.eval_all();
657 }
658
659 #[inline]
661 pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
662 self.core.guard_ref_slot(slot);
663 self.set_inputs(coords);
664 if !self.force_run
665 && slot < self.slot_provenance.len()
666 && !self.slot_provenance[slot].intersects(&self.changed_mask)
667 {
668 return self.core.buffer[slot];
669 }
670 self.force_run = false;
671 self.core.drive.stale = true;
672 self.core.eval_all();
673 self.core.buffer[slot]
674 }
675
676 pub fn engine_counts(&self) -> (usize, usize) {
679 self.core.engine_counts()
680 }
681
682 pub fn retain_nodes(&mut self, nodes: Vec<Box<dyn PolydatNode>>) {
684 self.core._nodes = std::sync::Arc::new(nodes);
685 }
686}
687
688pub type HybridKernel = HybridKernelPushPull;
693
694fn flatten_input_slots(
698 wiring: &[Vec<WireSource>],
699 nodes: &[Box<dyn PolydatNode>],
700 node_idx: usize,
701 port_offsets: &[Vec<usize>],
702 input_starts: &[usize],
703 input_widths: &[usize],
704) -> Vec<usize> {
705 let mut slots = Vec::new();
706 for source in &wiring[node_idx] {
707 let (start, w) = match source {
708 WireSource::Input(c) => (
709 input_starts.get(*c).copied().unwrap_or(*c),
710 input_widths.get(*c).copied().unwrap_or(1),
711 ),
712 WireSource::NodeOutput(u, p) => (
713 port_offsets[*u][*p],
714 nodes[*u].meta().outs[*p].typ.slot_width(),
715 ),
716 };
717 slots.extend(start..start + w);
718 }
719 slots
720}
721
722fn flatten_ref_output_starts(
725 nodes: &[Box<dyn PolydatNode>],
726 node_idx: usize,
727 port_offsets: &[Vec<usize>],
728) -> Vec<usize> {
729 nodes[node_idx]
730 .meta()
731 .outs
732 .iter()
733 .enumerate()
734 .filter(|(_, out)| out.typ.slot_color() == crate::ast::SlotColor::Ref2)
735 .map(|(p, _)| port_offsets[node_idx][p])
736 .collect()
737}
738
739fn flatten_output_slots(
741 nodes: &[Box<dyn PolydatNode>],
742 node_idx: usize,
743 port_offsets: &[Vec<usize>],
744) -> Vec<usize> {
745 let mut slots = Vec::new();
746 for (p, out) in nodes[node_idx].meta().outs.iter().enumerate() {
747 let start = port_offsets[node_idx][p];
748 slots.extend(start..start + out.typ.slot_width());
749 }
750 slots
751}
752
753fn refused(reason: String) -> crate::KernelError {
764 crate::KernelError::Refused {
765 engine: crate::compile::select::Engine::Native(crate::compile::select::Provenance::Auto),
766 reason,
767 }
768}
769
770#[cfg(feature = "jit")]
772#[allow(clippy::too_many_arguments)]
773pub(crate) fn build_hybrid(
774 nodes: &[Box<dyn PolydatNode>],
775 wiring: &[Vec<WireSource>],
776 coord_count: usize,
777 total_slots: usize,
778 port_offsets: &[Vec<usize>],
779 input_starts: &[usize],
780 input_widths: &[usize],
781 output_map: HashMap<String, usize>,
782 ref_slots: Vec<bool>,
783 input_types: &[crate::ast::PortType],
784 externs: crate::compile::externs::Externs,
785 constant: Vec<bool>,
786 volatile: Vec<bool>,
787 attribution: std::sync::Arc<crate::compile::Attribution>,
788) -> Result<HybridKernelPushPull, crate::KernelError> {
789 let mut steps: Vec<Option<HybridStep>> = Vec::new();
792 let mut pending: Vec<PendingSegment> = Vec::new();
793 let mut scratch: Vec<crate::ast::ScratchBuf> = Vec::new();
794 let mut ref_scratch: Vec<(usize, usize)> = Vec::new();
795 let mut max_inputs = 0usize;
796 let mut max_outputs = 0usize;
797 let graph = GraphView {
798 nodes,
799 wiring,
800 port_offsets,
801 input_types,
802 };
803
804 let classifications: Vec<(JitOp, Vec<usize>, Vec<usize>)> = nodes
806 .iter()
807 .enumerate()
808 .map(|(node_idx, node)| {
809 let wire_types: Vec<crate::ast::PortType> = wiring[node_idx]
813 .iter()
814 .map(|src| match src {
815 WireSource::Input(c) => input_types
816 .get(*c)
817 .copied()
818 .unwrap_or(crate::ast::PortType::U64),
819 WireSource::NodeOutput(j, p) => nodes[*j].meta().outs[*p].typ,
820 })
821 .collect();
822 let jit_op = jit::classify_node_typed(node.as_ref(), &wire_types);
823
824 let input_slots = flatten_input_slots(
825 wiring,
826 nodes,
827 node_idx,
828 port_offsets,
829 input_starts,
830 input_widths,
831 );
832 let output_slots = flatten_output_slots(nodes, node_idx, port_offsets);
833
834 max_inputs = max_inputs.max(input_slots.len());
835 max_outputs = max_outputs.max(output_slots.len());
836
837 (jit_op, input_slots, output_slots)
838 })
839 .collect();
840 let mut classifications = classifications;
850 let mut eligible = vec![false; nodes.len()];
851 for (node_idx, node) in nodes.iter().enumerate() {
852 if matches!(classifications[node_idx].0, JitOp::Fallback) {
853 continue;
854 }
855 if !crate::compile::none_rule_admits(
856 node.accepts_none_inputs(),
857 &wiring[node_idx],
858 &eligible,
859 ) {
860 classifications[node_idx].0 = JitOp::Fallback;
861 continue;
862 }
863 eligible[node_idx] = true;
864 }
865 let unset = externs.unset_slots();
871 if !unset.is_empty() {
872 let mut tainted = vec![false; nodes.len()];
873 for node_idx in 0..nodes.len() {
874 tainted[node_idx] = wiring[node_idx].iter().any(|src| match src {
875 WireSource::Input(c) => unset.contains(&input_starts[*c]),
876 WireSource::NodeOutput(j, _) => tainted[*j],
877 });
878 if tainted[node_idx] {
879 classifications[node_idx].0 = JitOp::Fallback;
880 }
881 }
882 }
883
884 let mut node_step = vec![usize::MAX; nodes.len()];
886 let order: Vec<usize> = (0..nodes.len())
894 .filter(|&k| constant[k])
895 .chain((0..nodes.len()).filter(|&k| !constant[k]))
896 .collect();
897 let mut rank = vec![0usize; nodes.len()];
898 for (pos, &k) in order.iter().enumerate() {
899 rank[k] = pos;
900 }
901 let is_side = |k: usize| matches!(nodes[k].purity(), crate::ast::Purity::SideChannel { .. });
910 let preds: Vec<Vec<usize>> = wiring
911 .iter()
912 .map(|w| {
913 w.iter()
914 .filter_map(|src| match src {
915 WireSource::NodeOutput(j, _) => Some(*j),
916 WireSource::Input(_) => None,
917 })
918 .collect()
919 })
920 .collect();
921 let fusible: Vec<bool> = (0..nodes.len())
922 .map(|k| !matches!(classifications[k].0, JitOp::Fallback) && !is_side(k))
923 .collect();
924 let class: Vec<u64> = (0..nodes.len())
925 .map(|k| constant[k] as u64 | (volatile[k] as u64) << 1)
926 .collect();
927 let plan =
928 crate::compile::fusion_units::plan_units(&preds, &fusible, &class, &rank, &|c| c & 1 == 1);
929 for members in plan.units {
930 let i = members[0];
931 if matches!(classifications[i].0, JitOp::Fallback) {
932 let (_, ref input_slots, ref output_slots) = classifications[i];
936 let step = closure_step_for(
937 &graph,
938 i,
939 input_slots.clone(),
940 output_slots.clone(),
941 &mut scratch,
942 &mut ref_scratch,
943 )?;
944 node_step[i] = steps.len();
945 steps.push(Some(HybridStep::Closure(step)));
946 } else {
947 for &k in &members {
951 let base = scratch.len();
952 classifications[k].0.place_scratch(base);
953 let elems = classifications[k].0.scratch_elems().to_vec();
954 ref_scratch.extend(crate::compile::assembly::scratch_pairs(
955 &nodes[k].meta().name,
956 &flatten_ref_output_starts(nodes, k, port_offsets),
957 &elems,
958 base,
959 ));
960 scratch.extend(elems.iter().map(|e| crate::ast::ScratchBuf::new(*e)));
961 }
962 let batch: Vec<(JitOp, Vec<usize>, Vec<usize>)> = members
970 .iter()
971 .map(|&k| classifications[k].clone())
972 .collect();
973 let written: std::collections::HashSet<usize> = batch
974 .iter()
975 .flat_map(|(_, _, o)| o.iter().copied())
976 .collect();
977 let mut input_slots: Vec<usize> = Vec::new();
978 for (_, ins, _) in &batch {
979 for &s in ins {
980 if !written.contains(&s) && !input_slots.contains(&s) {
981 input_slots.push(s);
982 }
983 }
984 }
985 let output_slots: Vec<usize> = batch
986 .iter()
987 .flat_map(|(_, _, o)| o.iter().copied())
988 .collect();
989 let segment = steps.len();
990 for &k in &members {
991 node_step[k] = segment;
992 }
993 steps.push(None);
994 pending.push(PendingSegment {
995 step: segment,
996 batch,
997 input_slots,
998 output_slots,
999 nodes: members,
1000 });
1001 }
1002 }
1003 let batches: Vec<&[jit::JitStep]> = pending.iter().map(|p| p.batch.as_slice()).collect();
1007 let (entries, code) = if batches.is_empty() {
1008 (Vec::new(), None)
1009 } else {
1010 let (entries, code) =
1011 jit::compile_jit_entries(&batches, Some(total_slots)).map_err(refused)?;
1012 (entries, Some(code))
1013 };
1014 for (p, (code_fn, fallible)) in pending.into_iter().zip(entries) {
1015 steps[p.step] = Some(HybridStep::Jit(JitSegment {
1016 code_fn,
1017 fallible,
1018 _module: code.clone().expect("a segment was compiled"),
1019 input_slots: p.input_slots,
1020 output_slots: p.output_slots,
1021 nodes: p.nodes,
1022 }));
1023 }
1024 let steps: Vec<HybridStep> = steps
1025 .into_iter()
1026 .map(|s| s.expect("every step is placed"))
1027 .collect();
1028
1029 let output_types = output_types_of(nodes, port_offsets, input_starts, input_types, &output_map);
1030 build_pushpull_from_steps(
1031 steps,
1032 scratch,
1033 ref_scratch,
1034 ref_slots,
1035 wiring,
1036 nodes,
1037 coord_count,
1038 total_slots,
1039 output_map,
1040 max_inputs,
1041 max_outputs,
1042 input_starts,
1043 input_widths,
1044 output_types,
1045 externs,
1046 constant,
1047 volatile,
1048 attribution,
1049 node_step,
1050 )
1051}
1052
1053fn output_types_of(
1056 nodes: &[Box<dyn PolydatNode>],
1057 port_offsets: &[Vec<usize>],
1058 input_starts: &[usize],
1059 input_types: &[crate::ast::PortType],
1060 output_map: &HashMap<String, usize>,
1061) -> HashMap<String, crate::ast::PortType> {
1062 let mut slot_types: HashMap<usize, crate::ast::PortType> = HashMap::new();
1063 for (start, ty) in input_starts.iter().zip(input_types) {
1064 slot_types.insert(*start, *ty);
1065 }
1066 for (node_idx, node) in nodes.iter().enumerate() {
1067 for (p, out) in node.meta().outs.iter().enumerate() {
1068 slot_types.insert(port_offsets[node_idx][p], out.typ);
1069 }
1070 }
1071 output_map
1072 .iter()
1073 .map(|(name, slot)| {
1074 (
1075 name.clone(),
1076 slot_types
1077 .get(slot)
1078 .copied()
1079 .unwrap_or(crate::ast::PortType::U64),
1080 )
1081 })
1082 .collect()
1083}
1084
1085#[cfg(not(feature = "jit"))]
1087#[allow(clippy::too_many_arguments)]
1088pub(crate) fn build_hybrid(
1089 nodes: &[Box<dyn PolydatNode>],
1090 wiring: &[Vec<WireSource>],
1091 coord_count: usize,
1092 total_slots: usize,
1093 port_offsets: &[Vec<usize>],
1094 input_starts: &[usize],
1095 input_widths: &[usize],
1096 output_map: HashMap<String, usize>,
1097 ref_slots: Vec<bool>,
1098 input_types: &[crate::ast::PortType],
1099 externs: crate::compile::externs::Externs,
1100 constant: Vec<bool>,
1101 volatile: Vec<bool>,
1102 attribution: std::sync::Arc<crate::compile::Attribution>,
1103) -> Result<HybridKernelPushPull, crate::KernelError> {
1104 let mut steps: Vec<HybridStep> = Vec::new();
1105 let mut scratch: Vec<crate::ast::ScratchBuf> = Vec::new();
1106 let mut ref_scratch: Vec<(usize, usize)> = Vec::new();
1107 let mut max_inputs = 0usize;
1108 let mut max_outputs = 0usize;
1109 let graph = GraphView {
1110 nodes,
1111 wiring,
1112 port_offsets,
1113 input_types,
1114 };
1115
1116 for node_idx in 0..nodes.len() {
1117 let input_slots = flatten_input_slots(
1118 wiring,
1119 nodes,
1120 node_idx,
1121 port_offsets,
1122 input_starts,
1123 input_widths,
1124 );
1125 let output_slots = flatten_output_slots(nodes, node_idx, port_offsets);
1126
1127 max_inputs = max_inputs.max(input_slots.len());
1128 max_outputs = max_outputs.max(output_slots.len());
1129
1130 let step = closure_step_for(
1131 &graph,
1132 node_idx,
1133 input_slots,
1134 output_slots,
1135 &mut scratch,
1136 &mut ref_scratch,
1137 )?;
1138 steps.push(HybridStep::Closure(step));
1139 }
1140 let node_step: Vec<usize> = (0..nodes.len()).collect();
1141
1142 let output_types = output_types_of(nodes, port_offsets, input_starts, input_types, &output_map);
1143 build_pushpull_from_steps(
1144 steps,
1145 scratch,
1146 ref_scratch,
1147 ref_slots,
1148 wiring,
1149 nodes,
1150 coord_count,
1151 total_slots,
1152 output_map,
1153 max_inputs,
1154 max_outputs,
1155 input_starts,
1156 input_widths,
1157 output_types,
1158 externs,
1159 constant,
1160 volatile,
1161 attribution,
1162 node_step,
1163 )
1164}
1165
1166#[derive(Clone, Copy)]
1171struct GraphView<'a> {
1172 nodes: &'a [Box<dyn PolydatNode>],
1173 wiring: &'a [Vec<WireSource>],
1174 port_offsets: &'a [Vec<usize>],
1175 input_types: &'a [crate::ast::PortType],
1176}
1177
1178fn closure_step_for(
1190 graph: &GraphView<'_>,
1191 node_idx: usize,
1192 input_slots: Vec<usize>,
1193 output_slots: Vec<usize>,
1194 scratch: &mut Vec<crate::ast::ScratchBuf>,
1195 ref_scratch: &mut Vec<(usize, usize)>,
1196) -> Result<ClosureStep, crate::KernelError> {
1197 let GraphView {
1198 nodes,
1199 wiring,
1200 port_offsets,
1201 input_types,
1202 } = *graph;
1203 let node = &nodes[node_idx];
1204 let scratch_start = scratch.len();
1205 let wire_types: Vec<crate::ast::PortType> = wiring[node_idx]
1206 .iter()
1207 .map(|src| match src {
1208 WireSource::Input(c) => input_types
1209 .get(*c)
1210 .copied()
1211 .unwrap_or(crate::ast::PortType::U64),
1212 WireSource::NodeOutput(j, p) => nodes[*j].meta().outs[*p].typ,
1213 })
1214 .collect();
1215 let op = if let Some(op) = node.compiled_u64() {
1216 ClosureOp::U64(op)
1217 } else if let Some(op) = crate::compile::assembly::identity_op(node.as_ref()) {
1218 ClosureOp::U64(op)
1219 } else if let Some(kit) = ref_copy_or_slot(node.as_ref(), &wire_types) {
1220 scratch.extend(kit.scratch.iter().map(|e| crate::ast::ScratchBuf::new(*e)));
1221 let starts = flatten_ref_output_starts(nodes, node_idx, port_offsets);
1222 ref_scratch.extend(crate::compile::assembly::scratch_pairs(
1223 &node.meta().name,
1224 &starts,
1225 &kit.scratch,
1226 scratch_start,
1227 ));
1228 ClosureOp::Slot(kit.op)
1229 } else {
1230 return Err(refused(format!(
1231 "node '{}' has no compiled form (docs/design/engines.md §8)",
1232 node.meta().name
1233 )));
1234 };
1235 Ok(ClosureStep {
1236 op,
1237 input_slots,
1238 output_slots,
1239 scratch_range: (scratch_start, scratch.len()),
1240 accepts_none: node.accepts_none_inputs(),
1241 node: node_idx,
1242 })
1243}
1244
1245#[allow(clippy::too_many_arguments)]
1251fn build_pushpull_from_steps(
1252 steps: Vec<HybridStep>,
1253 scratch: Vec<crate::ast::ScratchBuf>,
1254 ref_scratch: Vec<(usize, usize)>,
1255 ref_slots: Vec<bool>,
1256 wiring: &[Vec<WireSource>],
1257 nodes: &[Box<dyn PolydatNode>],
1258 coord_count: usize,
1259 total_slots: usize,
1260 output_map: HashMap<String, usize>,
1261 max_inputs: usize,
1262 max_outputs: usize,
1263 _input_starts: &[usize],
1264 input_widths: &[usize],
1265 output_types: HashMap<String, crate::ast::PortType>,
1266 externs: crate::compile::externs::Externs,
1267 constant: Vec<bool>,
1268 volatile: Vec<bool>,
1269 attribution: std::sync::Arc<crate::compile::Attribution>,
1270 node_step: Vec<usize>,
1271) -> Result<HybridKernelPushPull, crate::KernelError> {
1272 let step_count = steps.len();
1273 debug_assert_eq!(node_step.len(), nodes.len());
1274 debug_assert!(node_step.iter().all(|&s| s < step_count));
1275 let to_steps = |list: &[usize]| -> Vec<usize> {
1278 let mut v: Vec<usize> = list.iter().map(|&n| node_step[n]).collect();
1279 v.sort_unstable();
1280 v.dedup();
1281 v
1282 };
1283 let mut buffer = vec![0u64; total_slots + 1];
1285 let mut none = vec![false; total_slots];
1286 let any_none = externs.seed(&mut buffer, Some(&mut none));
1287
1288 let node_provenance = crate::kernel::PolydatProgram::compute_provenance(nodes, wiring);
1295 let input_dependents: Vec<Vec<usize>> =
1296 crate::kernel::PolydatProgram::compute_dependents(&node_provenance, input_widths.len())
1297 .iter()
1298 .map(|d| to_steps(d))
1299 .collect();
1300 let step_dependents: Vec<Vec<usize>> = input_widths
1301 .iter()
1302 .enumerate()
1303 .flat_map(|(i, w)| {
1304 std::iter::repeat_n(input_dependents.get(i).cloned().unwrap_or_default(), *w)
1305 })
1306 .collect();
1307
1308 let step_outs: Vec<&[usize]> = steps.iter().map(|s| s.output_slots()).collect();
1309 let slot_provenance =
1310 crate::compile::slot_provenance(coord_count, total_slots, &step_outs, &step_dependents);
1311
1312 debug_assert_eq!(constant.len(), nodes.len());
1317 debug_assert_eq!(volatile.len(), nodes.len());
1318 let mut step_constant = vec![true; step_count];
1319 let mut step_volatile = vec![false; step_count];
1320 let mut side = vec![false; step_count];
1321 for (n, node) in nodes.iter().enumerate() {
1322 let s = node_step[n];
1323 step_constant[s] &= constant[n];
1324 step_volatile[s] |= volatile[n];
1325 side[s] |= matches!(node.purity(), crate::ast::Purity::SideChannel { .. });
1326 }
1327 let volatile = step_volatile;
1328 let constants: Vec<usize> = (0..step_count).filter(|&i| step_constant[i]).collect();
1329 let step_inputs: Vec<&[usize]> = steps.iter().map(|s| s.input_slots()).collect();
1330 let step_outputs: Vec<&[usize]> = steps.iter().map(|s| s.output_slots()).collect();
1331 let plan = crate::compile::Invalidation::from_provenance(
1332 step_dependents.clone(),
1333 &step_inputs,
1334 &step_outputs,
1335 &output_map,
1336 total_slots,
1337 );
1338 let mut slot_step: Vec<Option<usize>> = vec![None; total_slots];
1339 for (i, outs) in step_outputs.iter().enumerate() {
1340 for &s in outs.iter() {
1341 slot_step[s] = Some(i);
1342 }
1343 }
1344 drop(step_inputs);
1345 drop(step_outputs);
1346
1347 let dirty: Vec<Vec<usize>> = plan.input_dependents.clone();
1348 let volatile_steps: Vec<usize> = (0..step_count).filter(|&i| volatile[i]).collect();
1349 let mut kernel = HybridKernelPushPull {
1350 core: HybridCore {
1351 engine: Engine::Native(Provenance::PushPull),
1352 buffer,
1353 coord_count,
1354 steps: std::sync::Arc::new(steps),
1355 output_map,
1356 gather_buf: vec![0u64; max_inputs.max(1)],
1357 scatter_buf: vec![0u64; max_outputs.max(1)],
1358 scratch,
1359 ref_slots,
1360 ref_scratch,
1361 output_types,
1362 externs,
1363 traversals: Vec::new().into(),
1364 resolved_outputs: Vec::new(),
1365 _nodes: std::sync::Arc::new(Vec::new()),
1366 drive: crate::compile::Drive {
1367 coords: Vec::new(),
1368 stale: true,
1369 },
1370 none,
1371 ran: vec![0; step_count],
1372 epoch: 0,
1373 all_ran: false,
1374 clean: vec![false; step_count],
1375 use_clean: true,
1376 plan: std::sync::Arc::new(plan),
1377 volatile: volatile.into(),
1378 side: side.into(),
1379 slot_step: slot_step.into(),
1380 sites: attribution,
1381 cur_step: 0,
1382 tracker: total_slots,
1383 all: (0..step_count).collect::<Vec<usize>>().into(),
1384 dirty: dirty.into(),
1385 any_none,
1386 volatile_steps: volatile_steps.into(),
1387 },
1388 slot_provenance,
1389 changed_mask: crate::kernel::ProvMask::all_below(coord_count), force_run: false,
1391 };
1392 kernel.core.begin_epoch();
1397 kernel.core.fold_steps(&constants)?;
1398 kernel.core.drive.stale = true;
1399 Ok(kernel)
1400}
1401
1402fn ref_copy_or_slot(
1406 node: &dyn PolydatNode,
1407 wire_types: &[crate::ast::PortType],
1408) -> Option<crate::ast::CompiledSlotKit> {
1409 let meta = node.meta();
1410 if (meta.name == "identity" || meta.name.starts_with("__port_"))
1411 && meta.outs.len() == 1
1412 && meta.outs[0].typ.slot_color() == crate::ast::SlotColor::Ref2
1413 {
1414 return crate::compile::assembly::ref_copy_kit(meta.outs[0].typ);
1415 }
1416 node.compiled_slot(
1417 wire_types,
1418 crate::compile::select::Engine::Native(crate::compile::select::Provenance::Auto),
1419 )
1420}
1421
1422impl HybridKernelRaw {
1425 fn mark_all_dirty(&mut self) {}
1427}
1428
1429impl HybridKernelPull {
1430 fn mark_all_dirty(&mut self) {
1432 self.changed_mask = crate::kernel::ProvMask::all_below(self.core.coord_count);
1433 self.force_run = true;
1434 }
1435}
1436
1437impl HybridKernelPushPull {
1438 fn mark_all_dirty(&mut self) {
1440 self.core.clean.fill(false);
1441 self.changed_mask = crate::kernel::ProvMask::all_below(self.core.coord_count);
1442 self.force_run = true;
1443 }
1444
1445 pub(crate) fn into_raw(self) -> HybridKernelRaw {
1448 let mut core = self.core;
1449 core.set_use_clean(false);
1450 core.engine = Engine::Native(Provenance::Raw);
1451 HybridKernelRaw { core }
1452 }
1453
1454 pub(crate) fn into_pull(self) -> HybridKernelPull {
1455 let mut core = self.core;
1456 core.set_use_clean(false);
1457 core.engine = Engine::Native(Provenance::Pull);
1458 let changed_mask = crate::kernel::ProvMask::all_below(core.coord_count);
1459 HybridKernelPull {
1460 core,
1461 slot_provenance: self.slot_provenance,
1462 changed_mask,
1463 force_run: false,
1464 }
1465 }
1466}
1467
1468use crate::compile::select::{Engine, Provenance};
1469
1470crate::compile::impl_kernel_trait!(HybridKernelRaw);
1471crate::compile::impl_kernel_trait!(HybridKernelPull);
1472crate::compile::impl_kernel_trait!(HybridKernelPushPull);
1473crate::compile::impl_slot_kernel!(HybridKernelRaw);
1474crate::compile::impl_slot_kernel!(HybridKernelPull);
1475crate::compile::impl_slot_kernel!(HybridKernelPushPull);
1476
1477#[inline(always)]
1492fn run_hybrid_step(
1493 step: &HybridStep,
1494 none_free: bool,
1495 buffer: &mut [u64],
1496 none: &mut [bool],
1497 gather: &mut [u64],
1498 scatter: &mut [u64],
1499 scratch: &mut [crate::ast::ScratchBuf],
1500) {
1501 if !none_free && !step.accepts_none() && step.input_slots().iter().any(|&s| none[s]) {
1502 for &s in step.output_slots() {
1503 none[s] = true;
1504 }
1505 return;
1506 }
1507 match step {
1508 #[cfg(feature = "jit")]
1509 HybridStep::Jit(seg) => {
1510 let code_fn = seg.code_fn;
1514 let buf_const = buffer.as_ptr();
1515 let buf_mut = buffer.as_mut_ptr();
1516 let sc = scratch.as_mut_ptr();
1517 if seg.fallible {
1518 crate::compile::jit::invoke_with_catch(move || unsafe {
1519 (code_fn)(buf_const, buf_mut, sc);
1520 });
1521 } else {
1522 unsafe { (code_fn)(buf_const, buf_mut, sc) };
1523 }
1524 }
1525 HybridStep::Closure(cs) => {
1526 for (i, &slot) in cs.input_slots.iter().enumerate() {
1527 gather[i] = buffer[slot];
1528 }
1529 match &cs.op {
1530 ClosureOp::U64(op) => op(
1531 &gather[..cs.input_slots.len()],
1532 &mut scatter[..cs.output_slots.len()],
1533 ),
1534 ClosureOp::Slot(op) => op(
1535 &gather[..cs.input_slots.len()],
1536 &mut scatter[..cs.output_slots.len()],
1537 &mut scratch[cs.scratch_range.0..cs.scratch_range.1],
1538 ),
1539 }
1540 for (i, &slot) in cs.output_slots.iter().enumerate() {
1541 buffer[slot] = scatter[i];
1542 }
1543 }
1544 }
1545 if !none_free {
1546 for &s in step.output_slots() {
1547 none[s] = false;
1548 }
1549 }
1550}