1use std::collections::{HashMap, HashSet, VecDeque};
8
9use serde::Serialize;
10use syn::spanned::Spanned;
11
12use crate::ir::SrcSpan;
13use crate::load::{self, Crate, ImplFn};
14use crate::op_semantics::{
15 self, affine_domain, domain_includes_zero, domain_violation, library_body, AbstractDtype,
16 BodyAtom, DomainRequirement, DomainViolationConfidence, DtypeRule, GradFlow, LibraryBody,
17 NumericDomain, OpEffect,
18};
19
20macro_rules! id_type {
21 ($name:ident) => {
22 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize)]
23 pub struct $name(pub usize);
24 };
25}
26
27id_type!(NodeId);
28id_type!(EdgeId);
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize)]
32#[serde(rename_all = "snake_case")]
33pub enum GradState {
34 Trainable,
36 Frozen,
38 Differentiable,
40 Severed,
42 LayoutDependent,
44 Unknown,
46}
47
48#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
50#[serde(tag = "kind", rename_all = "snake_case")]
51pub enum NodeKind {
52 Param { name: String },
54 Local { name: String },
56 Call { callee: String },
58 Literal { text: String },
60 Phi,
62 Return,
64 Unknown { reason: String },
66}
67
68#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
70#[serde(rename_all = "snake_case")]
71pub enum EdgeKind {
72 Data,
74 Severing,
76 Control,
78}
79
80#[derive(Debug, Clone, Serialize)]
81pub struct ExprNode {
82 pub id: NodeId,
83 pub kind: NodeKind,
84 pub span: SrcSpan,
85 pub shape: Option<String>,
87 pub dtype: AbstractDtype,
88 pub grad: GradState,
89 pub domain: NumericDomain,
91 pub type_name: Option<String>,
93}
94
95#[derive(Debug, Clone, Serialize)]
96pub struct ExprEdge {
97 pub id: EdgeId,
98 pub from: NodeId,
99 pub to: NodeId,
100 pub kind: EdgeKind,
101 pub label: Option<String>,
103}
104
105#[derive(Debug, Clone, Serialize)]
106pub struct DtypeConflict {
107 pub edge_or_node: NodeId,
108 pub op: String,
109 pub left: AbstractDtype,
110 pub right: AbstractDtype,
111 pub span: SrcSpan,
112 pub message: String,
113}
114
115#[derive(Debug, Clone, Serialize)]
120pub struct DtypeRisk {
121 pub edge_or_node: NodeId,
122 pub op: String,
123 pub known: AbstractDtype,
124 pub span: SrcSpan,
125 pub message: String,
126}
127
128#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Default)]
130#[serde(rename_all = "snake_case")]
131pub enum NumericImpact {
132 TrainingLossNaN,
134 GradientPoison,
136 InferenceOutputRisk,
138 #[default]
140 LocalOnly,
141}
142
143impl NumericImpact {
144 pub fn label(self) -> &'static str {
145 match self {
146 Self::TrainingLossNaN => "training impact: loss can become NaN when values saturate",
147 Self::GradientPoison => {
148 "training impact: gradients can be poisoned by non-finite values"
149 }
150 Self::InferenceOutputRisk => "inference impact: forward outputs can become non-finite",
151 Self::LocalOnly => "local numeric hazard (not shown to reach loss or return)",
152 }
153 }
154
155 pub fn is_training_failure(self) -> bool {
156 matches!(self, Self::TrainingLossNaN | Self::GradientPoison)
157 }
158}
159
160#[derive(Debug, Clone, Serialize)]
162pub struct NumericDomainViolation {
163 pub edge_or_node: NodeId,
164 pub op: String,
165 pub requires: String,
166 pub producer_domain: String,
167 pub proven: bool,
168 pub impact: NumericImpact,
169 pub span: SrcSpan,
170 pub message: String,
171 #[serde(skip_serializing_if = "Option::is_none")]
173 pub library_cite: Option<String>,
174}
175
176#[derive(Debug, Clone, Serialize)]
178pub struct ZeroTimesInfinity {
179 pub edge_or_node: NodeId,
180 pub impact: NumericImpact,
181 pub span: SrcSpan,
182 pub message: String,
183 #[serde(skip_serializing_if = "Option::is_none")]
184 pub library_cite: Option<String>,
185}
186
187#[derive(Debug, Clone, Serialize)]
188pub struct DataflowDiagnostic {
189 pub span: SrcSpan,
190 pub message: String,
191}
192
193#[derive(Debug, Clone, Default, Serialize)]
194pub struct ExprGraph {
195 pub nodes: Vec<ExprNode>,
196 pub edges: Vec<ExprEdge>,
197 pub entry_return: Option<NodeId>,
199 pub loss_nodes: Vec<NodeId>,
201 pub param_nodes: Vec<NodeId>,
203 pub tensor_nodes: Vec<NodeId>,
205 pub dtype_conflicts: Vec<DtypeConflict>,
206 pub dtype_risks: Vec<DtypeRisk>,
207 pub numeric_domain_violations: Vec<NumericDomainViolation>,
208 pub zero_times_infinity: Vec<ZeroTimesInfinity>,
209 pub diagnostics: Vec<DataflowDiagnostic>,
210}
211
212impl ExprGraph {
213 pub fn node(&self, id: NodeId) -> &ExprNode {
214 &self.nodes[id.0]
215 }
216
217 #[allow(dead_code)]
218 pub fn edge(&self, id: EdgeId) -> &ExprEdge {
219 &self.edges[id.0]
220 }
221
222 pub fn dead_params(&self) -> Vec<NodeId> {
224 let reachable = self.nodes_reaching_losses();
225 self.param_nodes
226 .iter()
227 .copied()
228 .filter(|id| {
229 matches!(self.node(*id).grad, GradState::Trainable) && !reachable.contains(id)
230 })
231 .collect()
232 }
233
234 pub fn severing_edges(&self) -> Vec<EdgeId> {
236 self.edges
237 .iter()
238 .filter(|e| matches!(e.kind, EdgeKind::Severing))
239 .map(|e| e.id)
240 .collect()
241 }
242
243 pub fn dtype_conflicts(&self) -> &[DtypeConflict] {
245 &self.dtype_conflicts
246 }
247
248 pub fn dtype_risks(&self) -> &[DtypeRisk] {
250 &self.dtype_risks
251 }
252
253 pub fn paths_to(&self, from: NodeId, to: NodeId) -> Vec<Vec<NodeId>> {
256 const MAX_PATHS: usize = 64;
257 const MAX_DEPTH: usize = 128;
258 let mut adj: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
259 for e in &self.edges {
260 if matches!(e.kind, EdgeKind::Severing) {
261 continue;
262 }
263 adj.entry(e.from).or_default().push(e.to);
264 }
265 let mut out = Vec::new();
266 let mut stack = vec![from];
267 let mut visiting = HashSet::new();
268 visiting.insert(from);
269 self.dfs_paths(
270 from,
271 to,
272 &adj,
273 &mut stack,
274 &mut visiting,
275 &mut out,
276 MAX_PATHS,
277 MAX_DEPTH,
278 );
279 out
280 }
281
282 #[allow(clippy::too_many_arguments)]
283 fn dfs_paths(
284 &self,
285 cur: NodeId,
286 to: NodeId,
287 adj: &HashMap<NodeId, Vec<NodeId>>,
288 stack: &mut Vec<NodeId>,
289 visiting: &mut HashSet<NodeId>,
290 out: &mut Vec<Vec<NodeId>>,
291 max_paths: usize,
292 max_depth: usize,
293 ) {
294 if out.len() >= max_paths || stack.len() > max_depth {
295 return;
296 }
297 if cur == to {
298 out.push(stack.clone());
299 return;
300 }
301 let Some(nexts) = adj.get(&cur) else {
302 return;
303 };
304 for &n in nexts {
305 if !visiting.insert(n) {
306 continue;
307 }
308 stack.push(n);
309 self.dfs_paths(n, to, adj, stack, visiting, out, max_paths, max_depth);
310 stack.pop();
311 visiting.remove(&n);
312 }
313 }
314
315 fn nodes_reaching_losses(&self) -> HashSet<NodeId> {
317 let mut rev: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
319 for e in &self.edges {
320 if matches!(e.kind, EdgeKind::Severing) {
321 continue;
323 }
324 rev.entry(e.to).or_default().push(e.from);
325 }
326 let mut seen = HashSet::new();
327 let mut q = VecDeque::new();
328 for &loss in &self.loss_nodes {
329 if seen.insert(loss) {
330 q.push_back(loss);
331 }
332 }
333 while let Some(n) = q.pop_front() {
334 if let Some(preds) = rev.get(&n) {
335 for &p in preds {
336 if seen.insert(p) {
337 q.push_back(p);
338 }
339 }
340 }
341 }
342 seen
343 }
344
345 pub fn reachable_from(&self, start: NodeId) -> HashSet<NodeId> {
347 let mut adj: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
348 for e in &self.edges {
349 if matches!(e.kind, EdgeKind::Severing) {
350 continue;
351 }
352 adj.entry(e.from).or_default().push(e.to);
353 }
354 let mut seen = HashSet::new();
355 let mut q = VecDeque::new();
356 seen.insert(start);
357 q.push_back(start);
358 while let Some(n) = q.pop_front() {
359 if let Some(nexts) = adj.get(&n) {
360 for &next in nexts {
361 if seen.insert(next) {
362 q.push_back(next);
363 }
364 }
365 }
366 }
367 seen
368 }
369
370 pub fn classify_numeric_impacts(&mut self) {
372 let can_reach_loss = self.nodes_reaching_losses();
373 let loss_nodes: HashSet<NodeId> = self.loss_nodes.iter().copied().collect();
374 let trainable: Vec<NodeId> = self
375 .param_nodes
376 .iter()
377 .copied()
378 .filter(|id| matches!(self.node(*id).grad, GradState::Trainable))
379 .collect();
380 let can_reach_return = {
381 let mut set = HashSet::new();
382 if let Some(ret) = self.entry_return {
383 let mut rev: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
384 for e in &self.edges {
385 if matches!(e.kind, EdgeKind::Severing) {
386 continue;
387 }
388 rev.entry(e.to).or_default().push(e.from);
389 }
390 let mut q = VecDeque::new();
391 set.insert(ret);
392 q.push_back(ret);
393 while let Some(n) = q.pop_front() {
394 if let Some(preds) = rev.get(&n) {
395 for &p in preds {
396 if set.insert(p) {
397 q.push_back(p);
398 }
399 }
400 }
401 }
402 }
403 set
404 };
405
406 let from_trainable: HashSet<NodeId> = trainable
407 .iter()
408 .flat_map(|param| self.reachable_from(*param))
409 .collect();
410
411 let classify = |node: NodeId| -> NumericImpact {
412 let reaches_loss = loss_nodes.contains(&node) || can_reach_loss.contains(&node);
413 if reaches_loss {
414 return NumericImpact::TrainingLossNaN;
415 }
416 if from_trainable.contains(&node) {
417 return NumericImpact::GradientPoison;
418 }
419 if can_reach_return.contains(&node) {
420 return NumericImpact::InferenceOutputRisk;
421 }
422 NumericImpact::LocalOnly
423 };
424
425 let domain_nodes: Vec<NodeId> = self
426 .numeric_domain_violations
427 .iter()
428 .map(|finding| finding.edge_or_node)
429 .collect();
430 let domain_impacts: Vec<NumericImpact> =
431 domain_nodes.iter().copied().map(classify).collect();
432 for (finding, impact) in self
433 .numeric_domain_violations
434 .iter_mut()
435 .zip(domain_impacts)
436 {
437 finding.impact = impact;
438 if !finding.message.contains("impact:") {
439 finding
440 .message
441 .push_str(&format!(" ({})", finding.impact.label()));
442 }
443 }
444
445 let zti_nodes: Vec<NodeId> = self
446 .zero_times_infinity
447 .iter()
448 .map(|finding| finding.edge_or_node)
449 .collect();
450 let zti_impacts: Vec<NumericImpact> = zti_nodes.iter().copied().map(classify).collect();
451 for (finding, impact) in self.zero_times_infinity.iter_mut().zip(zti_impacts) {
452 finding.impact = impact;
453 if !finding.message.contains("impact:") {
454 finding
455 .message
456 .push_str(&format!(" ({})", finding.impact.label()));
457 }
458 }
459 }
460}
461
462pub fn analyze(krate: &Crate, entrypoint: &str) -> anyhow::Result<ExprGraph> {
467 analyze_with_candle_version(krate, entrypoint, None)
468}
469
470pub fn analyze_with_candle_version(
472 krate: &Crate,
473 entrypoint: &str,
474 candle_nn_version: Option<&str>,
475) -> anyhow::Result<ExprGraph> {
476 let mut analyzer = Analyzer::new(krate, candle_nn_version);
477 analyzer.run(entrypoint)?;
478 analyzer.graph.classify_numeric_impacts();
479 Ok(analyzer.graph)
480}
481
482struct Analyzer<'a> {
483 krate: &'a Crate,
484 graph: ExprGraph,
485 node_types: HashMap<NodeId, String>,
488 call_stack: HashSet<(String, String)>,
490 file_hint: usize,
491 module_path: String,
492 candle_nn_version: Option<String>,
493 expanding_library_body: bool,
495 expansion_cite: Option<&'static str>,
497}
498
499impl<'a> Analyzer<'a> {
500 fn new(krate: &'a Crate, candle_nn_version: Option<&str>) -> Self {
501 Self {
502 krate,
503 graph: ExprGraph::default(),
504 node_types: HashMap::new(),
505 call_stack: HashSet::new(),
506 file_hint: 0,
507 module_path: String::new(),
508 candle_nn_version: candle_nn_version.map(str::to_string),
509 expanding_library_body: false,
510 expansion_cite: None,
511 }
512 }
513
514 fn run(&mut self, entrypoint: &str) -> anyhow::Result<()> {
515 let (func, type_name) = resolve_entrypoint(self.krate, entrypoint)?;
516 self.file_hint = func.span.file;
517 self.module_path = func.module_path.clone();
518 let key = (type_name.clone(), func.fn_name.clone());
519 self.call_stack.insert(key.clone());
520 let ret = self.analyze_function(func, &type_name, None, &[])?;
521 self.call_stack.remove(&key);
522 self.graph.entry_return = ret;
523 self.finalize_tensor_nodes();
524 Ok(())
525 }
526
527 fn add_node(
528 &mut self,
529 kind: NodeKind,
530 span: SrcSpan,
531 dtype: AbstractDtype,
532 grad: GradState,
533 shape: Option<String>,
534 ) -> NodeId {
535 self.add_node_with_domain(kind, span, dtype, grad, shape, NumericDomain::Unknown)
536 }
537
538 fn add_node_with_domain(
539 &mut self,
540 kind: NodeKind,
541 span: SrcSpan,
542 dtype: AbstractDtype,
543 grad: GradState,
544 shape: Option<String>,
545 domain: NumericDomain,
546 ) -> NodeId {
547 let id = NodeId(self.graph.nodes.len());
548 self.graph.nodes.push(ExprNode {
549 id,
550 kind,
551 span,
552 shape,
553 dtype,
554 grad,
555 domain,
556 type_name: None,
557 });
558 id
559 }
560
561 fn add_edge(
562 &mut self,
563 from: NodeId,
564 to: NodeId,
565 kind: EdgeKind,
566 label: Option<String>,
567 ) -> EdgeId {
568 let id = EdgeId(self.graph.edges.len());
569 self.graph.edges.push(ExprEdge {
570 id,
571 from,
572 to,
573 kind,
574 label,
575 });
576 id
577 }
578
579 fn diagnose(&mut self, span: SrcSpan, message: impl Into<String>) {
580 self.graph.diagnostics.push(DataflowDiagnostic {
581 span,
582 message: message.into(),
583 });
584 }
585
586 fn set_node_type(&mut self, id: NodeId, ty: impl Into<String>) {
587 let ty = ty.into();
588 if !ty.is_empty() && ty != "()" {
589 self.node_types.insert(id, ty.clone());
590 self.graph.nodes[id.0].type_name = Some(ty);
591 }
592 }
593
594 fn node_type(&self, id: NodeId) -> Option<&str> {
595 self.node_types.get(&id).map(String::as_str)
596 }
597
598 fn finalize_tensor_nodes(&mut self) {
599 let mut tensor = self
600 .graph
601 .nodes
602 .iter()
603 .map(|node| {
604 node.type_name
605 .as_deref()
606 .is_some_and(is_candle_tensor_receiver)
607 })
608 .collect::<Vec<_>>();
609 loop {
610 let mut changed = false;
611 for node in &self.graph.nodes {
612 if tensor[node.id.0] || !matches!(node.kind, NodeKind::Phi) {
613 continue;
614 }
615 let inputs = self
616 .graph
617 .edges
618 .iter()
619 .filter(|edge| edge.to == node.id)
620 .map(|edge| edge.from)
621 .collect::<Vec<_>>();
622 if !inputs.is_empty() && inputs.iter().all(|input| tensor[input.0]) {
623 tensor[node.id.0] = true;
624 changed = true;
625 }
626 }
627 if !changed {
628 break;
629 }
630 }
631 self.graph.tensor_nodes = tensor
632 .into_iter()
633 .enumerate()
634 .filter_map(|(index, is_tensor)| is_tensor.then_some(NodeId(index)))
635 .collect();
636 }
637
638 fn analyze_function(
639 &mut self,
640 func: &ImplFn,
641 type_name: &str,
642 receiver: Option<NodeId>,
643 args: &[NodeId],
644 ) -> anyhow::Result<Option<NodeId>> {
645 let mut env: HashMap<String, NodeId> = HashMap::new();
646 let mut arg_i = 0usize;
647 let _ = type_name;
648
649 for (param_index, pname) in func.params.iter().enumerate() {
650 let param_type = func
651 .param_types
652 .get(param_index)
653 .map(String::as_str)
654 .unwrap_or_default();
655 if pname == "self" {
656 let id = if let Some(recv) = receiver {
657 recv
658 } else {
659 self.add_node(
660 NodeKind::Param {
661 name: "self".into(),
662 },
663 func.span,
664 AbstractDtype::Unknown,
665 GradState::Unknown,
666 None,
667 )
668 };
669 self.set_node_type(id, type_name);
670 env.insert("self".to_string(), id);
671 continue;
672 }
673
674 let id = if arg_i < args.len() {
675 let src = args[arg_i];
676 arg_i += 1;
677 let g = hint_param_grad(param_type, self.graph.node(src).grad);
680 if g != self.graph.node(src).grad {
681 self.graph.nodes[src.0].grad = g;
682 }
683 if matches!(
684 self.graph.node(src).grad,
685 GradState::Trainable | GradState::Frozen
686 ) && !self.graph.param_nodes.contains(&src)
687 {
688 self.graph.param_nodes.push(src);
689 }
690 src
691 } else {
692 let grad = hint_param_grad(param_type, GradState::Unknown);
693 let id = self.add_node(
694 NodeKind::Param {
695 name: pname.clone(),
696 },
697 func.span,
698 AbstractDtype::Unknown,
699 grad,
700 None,
701 );
702 if matches!(grad, GradState::Trainable | GradState::Frozen) {
703 self.graph.param_nodes.push(id);
704 }
705 id
706 };
707 if let Some(base) = source_type_base(param_type) {
708 self.set_node_type(id, base);
709 }
710 env.insert(pname.clone(), id);
711 }
712
713 let mut last: Option<NodeId> = None;
714 for stmt in &func.block.stmts {
715 last = self.stmt(&mut env, stmt)?;
716 }
717 Ok(last)
718 }
719
720 fn stmt(
721 &mut self,
722 env: &mut HashMap<String, NodeId>,
723 stmt: &syn::Stmt,
724 ) -> anyhow::Result<Option<NodeId>> {
725 match stmt {
726 syn::Stmt::Local(local) => {
727 let init = match &local.init {
728 Some(local_init) => self.expr(env, &local_init.expr)?,
729 None => {
730 let span = span_of(self.file_hint, local.pat.span());
731 self.add_node(
732 NodeKind::Unknown {
733 reason: "uninitialized local".into(),
734 },
735 span,
736 AbstractDtype::Unknown,
737 GradState::Unknown,
738 None,
739 )
740 }
741 };
742 bind_pat(env, &local.pat, init);
743 Ok(Some(init))
744 }
745 syn::Stmt::Expr(expr, _) => Ok(Some(self.expr(env, expr)?)),
746 syn::Stmt::Item(_) => Ok(None),
747 syn::Stmt::Macro(m) => {
748 let span = span_of(self.file_hint, m.mac.path.span());
749 self.diagnose(span, "macro statement not expanded; left Unknown");
750 Ok(Some(self.add_node(
751 NodeKind::Unknown {
752 reason: "macro".into(),
753 },
754 span,
755 AbstractDtype::Unknown,
756 GradState::Unknown,
757 None,
758 )))
759 }
760 }
761 }
762
763 fn expr(
764 &mut self,
765 env: &mut HashMap<String, NodeId>,
766 expr: &syn::Expr,
767 ) -> anyhow::Result<NodeId> {
768 match expr {
769 syn::Expr::Path(p) => self.expr_path(env, p),
770 syn::Expr::Lit(l) => {
771 let span = span_of(self.file_hint, l.lit.span());
772 Ok(self.add_node(
773 NodeKind::Literal {
774 text: lit_text(&l.lit),
775 },
776 span,
777 AbstractDtype::Unknown,
778 GradState::Frozen,
779 None,
780 ))
781 }
782 syn::Expr::Reference(r) => self.expr(env, &r.expr),
783 syn::Expr::Paren(p) => self.expr(env, &p.expr),
784 syn::Expr::Group(g) => self.expr(env, &g.expr),
785 syn::Expr::Try(t) => self.expr(env, &t.expr),
786 syn::Expr::Unary(u) => self.expr(env, &u.expr),
787 syn::Expr::Field(f) => {
788 let base = self.expr(env, &f.base)?;
789 let span = span_of(self.file_hint, f.member.span());
790 let name = match &f.member {
791 syn::Member::Named(id) => id.to_string(),
792 syn::Member::Unnamed(i) => i.index.to_string(),
793 };
794 let id = self.add_node(
796 NodeKind::Local {
797 name: format!(".{name}"),
798 },
799 span,
800 AbstractDtype::Unknown,
801 GradState::Unknown,
802 None,
803 );
804 self.add_edge(base, id, EdgeKind::Data, Some("field".into()));
805 if let Some(owner_type) = self.node_type(base).map(ToString::to_string) {
806 let candidates = self.krate.struct_candidates(&owner_type);
807 if let [owner] = candidates.as_slice() {
808 if let Some(field) = owner.fields.iter().find(|field| field.name == name) {
809 if !field.ty.base.is_empty() {
810 self.set_node_type(id, field.ty.base.clone());
811 }
812 }
813 }
814 }
815 Ok(id)
819 }
820 syn::Expr::MethodCall(m) => self.method_call(env, m),
821 syn::Expr::Call(c) => self.func_call(env, c),
822 syn::Expr::Binary(b) => self.binary(env, b),
823 syn::Expr::If(i) => self.expr_if(env, i),
824 syn::Expr::Match(m) => self.expr_match(env, m),
825 syn::Expr::ForLoop(loop_expr) => {
826 let iterator = self.expr(env, &loop_expr.expr)?;
829 let mut loop_env = env.clone();
830 bind_pat(&mut loop_env, &loop_expr.pat, iterator);
831 let last = self.expr_block(&mut loop_env, &loop_expr.body)?;
832 Ok(last.unwrap_or_else(|| {
833 self.add_node(
834 NodeKind::Literal { text: "()".into() },
835 span_of(self.file_hint, loop_expr.for_token.span),
836 AbstractDtype::Unknown,
837 GradState::Unknown,
838 None,
839 )
840 }))
841 }
842 syn::Expr::While(loop_expr) => {
843 let _condition = self.expr(env, &loop_expr.cond)?;
844 let last = self.expr_block(env, &loop_expr.body)?;
845 Ok(last.unwrap_or_else(|| {
846 self.add_node(
847 NodeKind::Literal { text: "()".into() },
848 span_of(self.file_hint, loop_expr.while_token.span),
849 AbstractDtype::Unknown,
850 GradState::Unknown,
851 None,
852 )
853 }))
854 }
855 syn::Expr::Loop(loop_expr) => {
856 let last = self.expr_block(env, &loop_expr.body)?;
857 Ok(last.unwrap_or_else(|| {
858 self.add_node(
859 NodeKind::Literal { text: "()".into() },
860 span_of(self.file_hint, loop_expr.loop_token.span),
861 AbstractDtype::Unknown,
862 GradState::Unknown,
863 None,
864 )
865 }))
866 }
867 syn::Expr::Block(b) => {
868 let last = self.expr_block(env, &b.block)?;
869 Ok(last.unwrap_or_else(|| {
870 let span = span_of(self.file_hint, b.block.brace_token.span.join());
871 self.add_node(
872 NodeKind::Literal { text: "()".into() },
873 span,
874 AbstractDtype::Unknown,
875 GradState::Frozen,
876 None,
877 )
878 }))
879 }
880 syn::Expr::Return(r) => match &r.expr {
881 Some(e) => self.expr(env, e),
882 None => {
883 let span = span_of(self.file_hint, r.return_token.span);
884 Ok(self.add_node(
885 NodeKind::Return,
886 span,
887 AbstractDtype::Unknown,
888 GradState::Unknown,
889 None,
890 ))
891 }
892 },
893 syn::Expr::Closure(c) => {
894 let mut nested = env.clone();
896 let last = match &*c.body {
897 syn::Expr::Block(b) => self.expr_block(&mut nested, &b.block)?,
898 other => Some(self.expr(&mut nested, other)?),
899 };
900 Ok(last.unwrap_or_else(|| {
901 let span = span_of(self.file_hint, c.or1_token.span);
902 self.add_node(
903 NodeKind::Unknown {
904 reason: "empty closure".into(),
905 },
906 span,
907 AbstractDtype::Unknown,
908 GradState::Unknown,
909 None,
910 )
911 }))
912 }
913 syn::Expr::Tuple(t) => {
914 let mut last = None;
915 for e in &t.elems {
916 last = Some(self.expr(env, e)?);
917 }
918 Ok(last.unwrap_or_else(|| {
919 let span = SrcSpan::UNKNOWN;
920 self.add_node(
921 NodeKind::Literal { text: "()".into() },
922 span,
923 AbstractDtype::Unknown,
924 GradState::Frozen,
925 None,
926 )
927 }))
928 }
929 syn::Expr::Macro(m) => {
930 let span = span_of(self.file_hint, m.mac.path.span());
931 self.diagnose(
933 span,
934 format!("macro {} not expanded", path_last(&m.mac.path)),
935 );
936 Ok(self.add_node(
937 NodeKind::Unknown {
938 reason: "macro".into(),
939 },
940 span,
941 AbstractDtype::Unknown,
942 GradState::Unknown,
943 None,
944 ))
945 }
946 syn::Expr::Await(a) => self.expr(env, &a.base),
947 syn::Expr::Assign(a) => {
948 let val = self.expr(env, &a.right)?;
949 if let syn::Expr::Path(p) = &*a.left {
950 if let Some(name) = path_ident(p) {
951 env.insert(name, val);
952 }
953 }
954 Ok(val)
955 }
956 other => {
957 let span = SrcSpan {
958 file: self.file_hint,
959 line: 0,
960 col: 0,
961 };
962 self.diagnose(
963 span,
964 format!(
965 "unsupported expression {}; left Unknown",
966 expr_kind_name(other)
967 ),
968 );
969 Ok(self.add_node(
970 NodeKind::Unknown {
971 reason: expr_kind_name(other).into(),
972 },
973 span,
974 AbstractDtype::Unknown,
975 GradState::Unknown,
976 None,
977 ))
978 }
979 }
980 }
981
982 fn expr_path(
983 &mut self,
984 env: &mut HashMap<String, NodeId>,
985 p: &syn::ExprPath,
986 ) -> anyhow::Result<NodeId> {
987 if let Some(name) = path_ident(p) {
988 if let Some(id) = env.get(&name) {
989 return Ok(*id);
990 }
991 }
992 let text = path_text(&p.path);
993 let span = span_of(self.file_hint, p.path.span());
994 let segments: Vec<_> = p.path.segments.iter().collect();
997 if segments.len() >= 2 && segments[segments.len() - 2].ident == "DType" {
998 let dtype = AbstractDtype::parse(&segments[segments.len() - 1].ident.to_string());
999 return Ok(self.add_node(
1000 NodeKind::Literal { text },
1001 span,
1002 dtype,
1003 GradState::Frozen,
1004 None,
1005 ));
1006 }
1007 Ok(self.add_node(
1008 NodeKind::Unknown {
1009 reason: format!("unresolved path {text}"),
1010 },
1011 span,
1012 AbstractDtype::Unknown,
1013 GradState::Unknown,
1014 None,
1015 ))
1016 }
1017
1018 fn expr_block(
1019 &mut self,
1020 env: &mut HashMap<String, NodeId>,
1021 block: &syn::Block,
1022 ) -> anyhow::Result<Option<NodeId>> {
1023 let mut nested = env.clone();
1024 let mut last = None;
1025 for stmt in &block.stmts {
1026 last = self.stmt(&mut nested, stmt)?;
1027 }
1028 for (k, v) in nested {
1030 if env.contains_key(&k) {
1031 env.insert(k, v);
1032 }
1033 }
1034 Ok(last)
1035 }
1036
1037 fn expr_if(
1038 &mut self,
1039 env: &mut HashMap<String, NodeId>,
1040 i: &syn::ExprIf,
1041 ) -> anyhow::Result<NodeId> {
1042 let _cond = self.expr(env, &i.cond)?;
1043 let original_env = env.clone();
1044 let mut then_env = env.clone();
1045 let then_v = self.expr_block(&mut then_env, &i.then_branch)?;
1046 let mut else_env = original_env.clone();
1047 let else_v = match &i.else_branch {
1048 Some((_, e)) => Some(self.expr(&mut else_env, e)?),
1049 None => None,
1050 };
1051 let span = span_of(self.file_hint, i.if_token.span);
1052 self.merge_branch_envs(env, &original_env, &[then_env, else_env], span);
1053 let phi = self.add_node(
1054 NodeKind::Phi,
1055 span,
1056 AbstractDtype::Unknown,
1057 GradState::Unknown,
1058 None,
1059 );
1060 let mut dtypes = Vec::new();
1061 let mut grads = Vec::new();
1062 if let Some(t) = then_v {
1063 self.add_edge(t, phi, EdgeKind::Control, Some("then".into()));
1064 dtypes.push(self.graph.node(t).dtype);
1065 grads.push(self.graph.node(t).grad);
1066 }
1067 if let Some(e) = else_v {
1068 self.add_edge(e, phi, EdgeKind::Control, Some("else".into()));
1069 dtypes.push(self.graph.node(e).dtype);
1070 grads.push(self.graph.node(e).grad);
1071 }
1072 self.graph.nodes[phi.0].dtype = join_dtypes(&dtypes);
1073 self.graph.nodes[phi.0].grad = join_grads(&grads);
1074 Ok(phi)
1075 }
1076
1077 fn merge_branch_envs(
1078 &mut self,
1079 env: &mut HashMap<String, NodeId>,
1080 original: &HashMap<String, NodeId>,
1081 branches: &[HashMap<String, NodeId>],
1082 span: SrcSpan,
1083 ) {
1084 for (name, original_id) in original {
1085 let values = branches
1086 .iter()
1087 .map(|branch| branch.get(name).copied().unwrap_or(*original_id))
1088 .collect::<Vec<_>>();
1089 if values.iter().all(|value| *value == values[0]) {
1090 env.insert(name.clone(), values[0]);
1091 continue;
1092 }
1093 let phi = self.add_node(
1094 NodeKind::Phi,
1095 span,
1096 join_dtypes(
1097 &values
1098 .iter()
1099 .map(|value| self.graph.node(*value).dtype)
1100 .collect::<Vec<_>>(),
1101 ),
1102 join_grads(
1103 &values
1104 .iter()
1105 .map(|value| self.graph.node(*value).grad)
1106 .collect::<Vec<_>>(),
1107 ),
1108 None,
1109 );
1110 for (index, value) in values.iter().enumerate() {
1111 self.add_edge(
1112 *value,
1113 phi,
1114 EdgeKind::Control,
1115 Some(format!("branch{index}")),
1116 );
1117 }
1118 let types = values
1119 .iter()
1120 .filter_map(|value| self.node_type(*value))
1121 .collect::<Vec<_>>();
1122 if let Some(first) = types.first().copied() {
1123 if types.len() == values.len() && types.iter().all(|value| *value == first) {
1124 self.set_node_type(phi, first.to_string());
1125 }
1126 }
1127 env.insert(name.clone(), phi);
1128 }
1129 }
1130
1131 fn expr_match(
1132 &mut self,
1133 env: &mut HashMap<String, NodeId>,
1134 m: &syn::ExprMatch,
1135 ) -> anyhow::Result<NodeId> {
1136 let _scrut = self.expr(env, &m.expr)?;
1137 let span = span_of(self.file_hint, m.match_token.span);
1138 let phi = self.add_node(
1139 NodeKind::Phi,
1140 span,
1141 AbstractDtype::Unknown,
1142 GradState::Unknown,
1143 None,
1144 );
1145 let mut dtypes = Vec::new();
1146 let mut grads = Vec::new();
1147 for arm in &m.arms {
1148 let mut arm_env = env.clone();
1149 let v = self.expr(&mut arm_env, &arm.body)?;
1150 self.add_edge(v, phi, EdgeKind::Control, Some("arm".into()));
1151 dtypes.push(self.graph.node(v).dtype);
1152 grads.push(self.graph.node(v).grad);
1153 }
1154 self.graph.nodes[phi.0].dtype = join_dtypes(&dtypes);
1155 self.graph.nodes[phi.0].grad = join_grads(&grads);
1156 Ok(phi)
1157 }
1158
1159 fn binary(
1160 &mut self,
1161 env: &mut HashMap<String, NodeId>,
1162 b: &syn::ExprBinary,
1163 ) -> anyhow::Result<NodeId> {
1164 let left = self.expr(env, &b.left)?;
1165 let right = self.expr(env, &b.right)?;
1166 let op = match b.op {
1167 syn::BinOp::Add(_) => "add",
1168 syn::BinOp::Sub(_) => "sub",
1169 syn::BinOp::Mul(_) => "mul",
1170 syn::BinOp::Div(_) => "div",
1171 _ => {
1172 let span = span_of(self.file_hint, span_binop(&b.op));
1173 let id = self.add_node(
1174 NodeKind::Call {
1175 callee: "binary".into(),
1176 },
1177 span,
1178 AbstractDtype::Unknown,
1179 GradState::Unknown,
1180 None,
1181 );
1182 self.add_edge(left, id, EdgeKind::Data, Some("lhs".into()));
1183 self.add_edge(right, id, EdgeKind::Data, Some("rhs".into()));
1184 return Ok(id);
1185 }
1186 };
1187 let span = span_of(self.file_hint, span_binop(&b.op));
1188 if is_scalar_literal(&b.left) ^ is_scalar_literal(&b.right) {
1189 let tensor = if is_scalar_literal(&b.left) {
1190 right
1191 } else {
1192 left
1193 };
1194 let id = self.add_node(
1195 NodeKind::Call {
1196 callee: op.to_string(),
1197 },
1198 span,
1199 self.graph.node(tensor).dtype,
1200 self.graph.node(tensor).grad,
1201 self.graph.node(tensor).shape.clone(),
1202 );
1203 self.set_node_type(id, "Tensor");
1204 self.add_edge(left, id, EdgeKind::Data, Some("lhs".into()));
1205 self.add_edge(right, id, EdgeKind::Data, Some("rhs".into()));
1206 return Ok(id);
1207 }
1208 Ok(self.apply_op(op, span, &[left, right], None))
1209 }
1210
1211 fn method_call(
1212 &mut self,
1213 env: &mut HashMap<String, NodeId>,
1214 m: &syn::ExprMethodCall,
1215 ) -> anyhow::Result<NodeId> {
1216 let receiver = self.expr(env, &m.receiver)?;
1217 let mut args = Vec::with_capacity(m.args.len());
1218 for a in &m.args {
1219 args.push(self.expr(env, a)?);
1220 }
1221 let method = m.method.to_string();
1222 let span = span_of(self.file_hint, m.method.span());
1223
1224 if let Some(ret) = self.try_interproc_method(&method, receiver, &args, span)? {
1228 return Ok(ret);
1229 }
1230
1231 if let Some(receiver_type) = self.node_type(receiver).map(str::to_string) {
1232 let effect = op_semantics::lookup_method(
1233 &receiver_type,
1234 &method,
1235 self.candle_nn_version.as_deref(),
1236 );
1237 if !matches!(effect.dtype, DtypeRule::Unknown)
1238 || !matches!(effect.grad, GradFlow::Unknown)
1239 {
1240 let label = effect.name.clone();
1241 let result = self.apply_effect(&label, effect, span, &args, None);
1242 if !is_candle_tensor_receiver(&receiver_type) {
1243 self.add_edge(receiver, result, EdgeKind::Data, Some("module".into()));
1244 }
1245 return Ok(result);
1246 }
1247 }
1248
1249 let mut operands = vec![receiver];
1250 operands.extend(args.iter().copied());
1251 if !self
1252 .node_type(receiver)
1253 .is_some_and(is_candle_tensor_receiver)
1254 {
1255 let id = self.add_node(
1256 NodeKind::Call {
1257 callee: method.clone(),
1258 },
1259 span,
1260 AbstractDtype::Unknown,
1261 GradState::Unknown,
1262 None,
1263 );
1264 self.add_edge(receiver, id, EdgeKind::Data, Some("self".into()));
1265 for (index, arg) in args.iter().enumerate() {
1266 self.add_edge(*arg, id, EdgeKind::Data, Some(format!("arg{index}")));
1267 }
1268 self.diagnose(
1269 span,
1270 format!(
1271 "receiver type for method `{method}` is not proven to be Tensor or a \
1272 crate-local type; transfer semantics left unknown"
1273 ),
1274 );
1275 return Ok(id);
1276 }
1277 let explicit = if method == "to_dtype" {
1278 args.first().map(|id| self.graph.node(*id).dtype)
1279 } else {
1280 None
1281 };
1282 Ok(self.apply_op(&method, span, &operands, explicit))
1283 }
1284
1285 fn try_interproc_method(
1286 &mut self,
1287 method: &str,
1288 receiver: NodeId,
1289 args: &[NodeId],
1290 span: SrcSpan,
1291 ) -> anyhow::Result<Option<NodeId>> {
1292 let Some(receiver_type) = self.node_type(receiver).map(ToString::to_string) else {
1293 let count = self
1294 .krate
1295 .all_methods()
1296 .filter(|func| func.fn_name == method)
1297 .count();
1298 if count > 0 {
1299 self.diagnose(
1300 span,
1301 format!(
1302 "method `{method}` has {count} crate-local candidate(s), but the receiver \
1303 type is unknown; not inlined"
1304 ),
1305 );
1306 }
1307 return Ok(None);
1308 };
1309 if is_candle_tensor_receiver(&receiver_type) {
1310 return Ok(None);
1311 }
1312
1313 let candidates = self.krate.method_candidates(&receiver_type, method);
1314 let func = match candidates.as_slice() {
1315 [] => return Ok(None),
1316 [func] => *func,
1317 _ => {
1318 self.diagnose(
1319 span,
1320 format!(
1321 "ambiguous method `{receiver_type}::{method}` ({} candidates); not inlined",
1322 candidates.len()
1323 ),
1324 );
1325 return Ok(None);
1326 }
1327 };
1328
1329 let key = (func.type_name.clone(), func.fn_name.clone());
1330 if self.call_stack.contains(&key) {
1331 self.diagnose(
1332 span,
1333 format!("recursion guard: {}::{} already on stack", key.0, key.1),
1334 );
1335 let id = self.add_node(
1336 NodeKind::Call {
1337 callee: format!("{}::{method}", func.type_name),
1338 },
1339 span,
1340 AbstractDtype::Unknown,
1341 GradState::Unknown,
1342 None,
1343 );
1344 self.add_edge(receiver, id, EdgeKind::Data, Some("self".into()));
1345 for (i, a) in args.iter().enumerate() {
1346 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1347 }
1348 return Ok(Some(id));
1349 }
1350
1351 self.call_stack.insert(key.clone());
1352 let prev_file = self.file_hint;
1353 let prev_module = self.module_path.clone();
1354 self.file_hint = func.span.file;
1355 self.module_path = func.module_path.clone();
1356 let ret = self.analyze_function(func, &func.type_name, Some(receiver), args)?;
1357 self.file_hint = prev_file;
1358 self.module_path = prev_module;
1359 self.call_stack.remove(&key);
1360
1361 let out = ret.unwrap_or_else(|| {
1362 self.add_node(
1363 NodeKind::Call {
1364 callee: format!("{}::{method}", func.type_name),
1365 },
1366 span,
1367 AbstractDtype::Unknown,
1368 GradState::Unknown,
1369 None,
1370 )
1371 });
1372 let call_node = self.add_node(
1375 NodeKind::Call {
1376 callee: format!("{}::{method}", func.type_name),
1377 },
1378 span,
1379 self.graph.node(out).dtype,
1380 self.graph.node(out).grad,
1381 self.graph.node(out).shape.clone(),
1382 );
1383 if let Some(return_type) = resolved_return_type(func, &func.type_name) {
1384 self.set_node_type(call_node, return_type);
1385 }
1386 self.add_edge(out, call_node, EdgeKind::Data, Some("return".into()));
1387 Ok(Some(call_node))
1388 }
1389
1390 fn func_call(
1391 &mut self,
1392 env: &mut HashMap<String, NodeId>,
1393 c: &syn::ExprCall,
1394 ) -> anyhow::Result<NodeId> {
1395 let mut args = Vec::with_capacity(c.args.len());
1396 for a in &c.args {
1397 args.push(self.expr(env, a)?);
1398 }
1399 let span = match &*c.func {
1400 syn::Expr::Path(p) => span_of(self.file_hint, p.path.span()),
1401 _ => SrcSpan {
1402 file: self.file_hint,
1403 line: 0,
1404 col: 0,
1405 },
1406 };
1407
1408 if let syn::Expr::Path(p) = &*c.func {
1410 let source_segments: Vec<String> = p
1411 .path
1412 .segments
1413 .iter()
1414 .map(|s| s.ident.to_string())
1415 .collect();
1416 let segs = self
1417 .krate
1418 .resolve_import_path(&self.module_path, &source_segments);
1419 let last = segs.last().cloned().unwrap_or_default();
1420
1421 if segs.len() >= 2 {
1423 let type_name = normalize_qualified_segments(&segs[..segs.len() - 1]);
1424 let candidates = self.krate.method_candidates(&type_name, &last);
1425 match candidates.as_slice() {
1426 [func] => return self.call_crate_fn(func, &type_name, None, &args, span),
1427 [] => {}
1428 _ => {
1429 self.diagnose(
1430 span,
1431 format!(
1432 "call `{type_name}::{last}` is ambiguous ({} definitions); not \
1433 inlined",
1434 candidates.len()
1435 ),
1436 );
1437 }
1438 }
1439 }
1440
1441 let function_name = normalize_qualified_segments(&segs);
1443 let candidates = self.krate.function_candidates(&function_name);
1444 if let [func] = candidates.as_slice() {
1445 if op_semantics::lookup_for(&last, self.candle_nn_version.as_deref())
1448 .note
1449 .is_some()
1450 || matches!(
1451 op_semantics::lookup_for(&last, self.candle_nn_version.as_deref()).dtype,
1452 DtypeRule::SameAsInputs
1453 | DtypeRule::Preserve
1454 | DtypeRule::Explicit
1455 | DtypeRule::Fixed(_)
1456 )
1457 {
1458 }
1461 return self.call_crate_fn(func, "", None, &args, span);
1462 } else if candidates.len() > 1 {
1463 self.diagnose(
1464 span,
1465 format!(
1466 "free function `{function_name}` is ambiguous ({} definitions); not inlined",
1467 candidates.len()
1468 ),
1469 );
1470 }
1471
1472 let effect = op_semantics::lookup_for(&last, self.candle_nn_version.as_deref());
1474 if is_explicit_candle_path(&segs)
1475 && (effect.note.is_some()
1476 || !matches!(effect.grad, GradFlow::Unknown)
1477 || !matches!(effect.dtype, DtypeRule::Unknown))
1478 {
1479 let explicit = if last == "to_dtype" {
1480 args.first().map(|id| self.graph.node(*id).dtype)
1481 } else {
1482 None
1483 };
1484 return Ok(self.apply_op(&last, span, &args, explicit));
1485 }
1486
1487 let id = self.add_node(
1489 NodeKind::Call {
1490 callee: segs.join("::"),
1491 },
1492 span,
1493 AbstractDtype::Unknown,
1494 GradState::Unknown,
1495 None,
1496 );
1497 for (i, a) in args.iter().enumerate() {
1498 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1499 }
1500 return Ok(id);
1501 }
1502
1503 let callee = self.expr(env, &c.func)?;
1504 let id = self.add_node(
1505 NodeKind::Call {
1506 callee: "call".into(),
1507 },
1508 span,
1509 AbstractDtype::Unknown,
1510 GradState::Unknown,
1511 None,
1512 );
1513 self.add_edge(callee, id, EdgeKind::Data, Some("callee".into()));
1514 for (i, a) in args.iter().enumerate() {
1515 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1516 }
1517 Ok(id)
1518 }
1519
1520 fn call_crate_fn(
1521 &mut self,
1522 func: &ImplFn,
1523 type_name: &str,
1524 receiver: Option<NodeId>,
1525 args: &[NodeId],
1526 span: SrcSpan,
1527 ) -> anyhow::Result<NodeId> {
1528 let key = (func.type_name.clone(), func.fn_name.clone());
1529 let label = if type_name.is_empty() {
1530 func.fn_name.clone()
1531 } else {
1532 format!("{type_name}::{}", func.fn_name)
1533 };
1534
1535 if self.call_stack.contains(&key) {
1536 self.diagnose(span, format!("recursion guard: {label} already on stack"));
1537 let id = self.add_node(
1538 NodeKind::Call { callee: label },
1539 span,
1540 AbstractDtype::Unknown,
1541 GradState::Unknown,
1542 None,
1543 );
1544 for (i, a) in args.iter().enumerate() {
1545 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1546 }
1547 return Ok(id);
1548 }
1549
1550 self.call_stack.insert(key.clone());
1553 let prev_file = self.file_hint;
1554 let prev_module = self.module_path.clone();
1555 self.file_hint = func.span.file;
1556 self.module_path = func.module_path.clone();
1557 let ret = self.analyze_function(func, &func.type_name, receiver, args)?;
1558 self.file_hint = prev_file;
1559 self.module_path = prev_module;
1560 self.call_stack.remove(&key);
1561
1562 let out = ret.unwrap_or_else(|| {
1563 self.add_node(
1564 NodeKind::Call {
1565 callee: label.clone(),
1566 },
1567 span,
1568 AbstractDtype::Unknown,
1569 GradState::Unknown,
1570 None,
1571 )
1572 });
1573 let call_node = self.add_node(
1574 NodeKind::Call { callee: label },
1575 span,
1576 self.graph.node(out).dtype,
1577 self.graph.node(out).grad,
1578 self.graph.node(out).shape.clone(),
1579 );
1580 if let Some(return_type) = resolved_return_type(func, type_name) {
1581 self.set_node_type(call_node, return_type);
1582 }
1583 self.add_edge(out, call_node, EdgeKind::Data, Some("return".into()));
1584 Ok(call_node)
1585 }
1586
1587 fn apply_op(
1589 &mut self,
1590 op: &str,
1591 span: SrcSpan,
1592 operands: &[NodeId],
1593 explicit_dtype: Option<AbstractDtype>,
1594 ) -> NodeId {
1595 let effect = op_semantics::lookup_for(op, self.candle_nn_version.as_deref());
1596 self.apply_effect(op, effect, span, operands, explicit_dtype)
1597 }
1598
1599 fn apply_effect(
1600 &mut self,
1601 op: &str,
1602 effect: OpEffect,
1603 span: SrcSpan,
1604 operands: &[NodeId],
1605 explicit_dtype: Option<AbstractDtype>,
1606 ) -> NodeId {
1607 let (dtype, grad, edge_kind) = transfer(&effect, operands, explicit_dtype, |id| {
1608 let n = &self.graph.nodes[id.0];
1609 (n.dtype, n.grad)
1610 });
1611 let domain = self.result_domain(&effect, op, operands);
1612
1613 if matches!(effect.dtype, DtypeRule::SameAsInputs) {
1614 let operand_dtypes: Vec<AbstractDtype> = operands
1615 .iter()
1616 .map(|id| self.graph.node(*id).dtype)
1617 .collect();
1618 let known: Vec<AbstractDtype> = operand_dtypes
1619 .iter()
1620 .copied()
1621 .filter(|d| d.is_known())
1622 .collect();
1623 if let Some((a, b)) = first_dtype_mismatch(&known) {
1624 let id = self.add_node_with_domain(
1625 NodeKind::Call {
1626 callee: op.to_string(),
1627 },
1628 span,
1629 AbstractDtype::Unknown, grad,
1631 None,
1632 domain,
1633 );
1634 self.graph.dtype_conflicts.push(DtypeConflict {
1635 edge_or_node: id,
1636 op: op.to_string(),
1637 left: a,
1638 right: b,
1639 span,
1640 message: format!("{op} requires same dtype, got {a} vs {b}"),
1641 });
1642 for (i, &src) in operands.iter().enumerate() {
1643 let label = operand_label(op, i, operands.len());
1644 self.add_edge(src, id, edge_kind, Some(label));
1645 }
1646 if effect.is_loss {
1647 self.graph.loss_nodes.push(id);
1648 }
1649 self.set_node_type(id, "Tensor");
1650 self.finish_numeric_call(id, op, &effect, operands, span);
1651 return id;
1652 }
1653 if let Some(known_dtype) = known.first().copied().filter(|_| {
1654 operand_dtypes
1655 .iter()
1656 .any(|dtype| !matches!(dtype, AbstractDtype::Unknown))
1657 && operand_dtypes
1658 .iter()
1659 .any(|dtype| matches!(dtype, AbstractDtype::Unknown))
1660 }) {
1661 let id = self.add_node_with_domain(
1662 NodeKind::Call {
1663 callee: op.to_string(),
1664 },
1665 span,
1666 dtype,
1667 grad,
1668 None,
1669 domain,
1670 );
1671 self.graph.dtype_risks.push(DtypeRisk {
1672 edge_or_node: id,
1673 op: op.to_string(),
1674 known: known_dtype,
1675 span,
1676 message: format!(
1677 "{op} requires matching dtypes; one operand is {known_dtype} and another is unknown"
1678 ),
1679 });
1680 for (i, &src) in operands.iter().enumerate() {
1681 let label = operand_label(op, i, operands.len());
1682 self.add_edge(src, id, edge_kind, Some(label));
1683 }
1684 if effect.is_loss {
1685 self.graph.loss_nodes.push(id);
1686 }
1687 self.set_node_type(id, "Tensor");
1688 self.finish_numeric_call(id, op, &effect, operands, span);
1689 return id;
1690 }
1691 }
1692
1693 let id = self.add_node_with_domain(
1694 NodeKind::Call {
1695 callee: op.to_string(),
1696 },
1697 span,
1698 dtype,
1699 grad,
1700 None,
1701 domain,
1702 );
1703 for (i, &src) in operands.iter().enumerate() {
1704 let label = operand_label(op, i, operands.len());
1705 self.add_edge(src, id, edge_kind, Some(label));
1706 }
1707 if effect.is_loss {
1708 self.graph.loss_nodes.push(id);
1709 }
1710 if !matches!(effect.dtype, DtypeRule::Unknown)
1711 || !matches!(effect.grad, GradFlow::Unknown)
1712 || operands.iter().any(|operand| {
1713 self.node_type(*operand)
1714 .is_some_and(is_candle_tensor_receiver)
1715 }) && !is_non_tensor_tensor_method(op)
1716 {
1717 self.set_node_type(id, "Tensor");
1718 }
1719 if let Some(note) = effect.note {
1720 if matches!(effect.grad, GradFlow::Unknown) {
1721 self.diagnose(span, note);
1722 }
1723 }
1724 self.finish_numeric_call(id, op, &effect, operands, span);
1725 id
1726 }
1727
1728 fn finish_numeric_call(
1729 &mut self,
1730 id: NodeId,
1731 op: &str,
1732 effect: &OpEffect,
1733 operands: &[NodeId],
1734 span: SrcSpan,
1735 ) {
1736 self.record_numeric_effects(id, op, effect, operands, span);
1737 if !self.expanding_library_body {
1738 if let Some(body) = library_body(op, self.candle_nn_version.as_deref()) {
1739 self.expand_library_body(id, body, operands, span);
1740 }
1741 }
1742 }
1743
1744 fn expand_library_body(
1746 &mut self,
1747 outer: NodeId,
1748 body: &LibraryBody,
1749 args: &[NodeId],
1750 span: SrcSpan,
1751 ) {
1752 self.expanding_library_body = true;
1753 self.expansion_cite = Some(body.cite);
1754 let mut vals: Vec<NodeId> = Vec::with_capacity(body.steps.len());
1755 for step in body.steps {
1756 let id = match *step {
1757 BodyAtom::Arg(index) => args.get(index).copied().unwrap_or_else(|| {
1758 self.add_node(
1759 NodeKind::Unknown {
1760 reason: format!("missing library body arg {index}"),
1761 },
1762 span,
1763 AbstractDtype::Unknown,
1764 GradState::Unknown,
1765 None,
1766 )
1767 }),
1768 BodyAtom::Assume { src, domain } => {
1769 let src = vals[src as usize];
1770 let id = self.add_node_with_domain(
1771 NodeKind::Call {
1772 callee: "library_domain_assume".into(),
1773 },
1774 span,
1775 self.graph.node(src).dtype,
1776 self.graph.node(src).grad,
1777 None,
1778 domain,
1779 );
1780 self.add_edge(src, id, EdgeKind::Data, Some("assume".into()));
1781 self.set_node_type(id, "Tensor");
1782 id
1783 }
1784 BodyAtom::Unary { op, src } => {
1785 let src = vals[src as usize];
1786 self.apply_op(op, span, &[src], None)
1787 }
1788 BodyAtom::Binary { op, left, right } => {
1789 let left = vals[left as usize];
1790 let right = vals[right as usize];
1791 self.apply_op(op, span, &[left, right], None)
1792 }
1793 BodyAtom::Affine { src, mul, add } => {
1794 let src = vals[src as usize];
1795 let mul_node = self.add_node(
1796 NodeKind::Literal {
1797 text: format!("{mul}"),
1798 },
1799 span,
1800 AbstractDtype::Unknown,
1801 GradState::Frozen,
1802 None,
1803 );
1804 let add_node = self.add_node(
1805 NodeKind::Literal {
1806 text: format!("{add}"),
1807 },
1808 span,
1809 AbstractDtype::Unknown,
1810 GradState::Frozen,
1811 None,
1812 );
1813 self.apply_op("affine", span, &[src, mul_node, add_node], None)
1814 }
1815 };
1816 vals.push(id);
1817 }
1818 if let Some(&last) = vals.last() {
1819 self.add_edge(last, outer, EdgeKind::Data, Some("expanded_body".into()));
1820 self.graph.nodes[outer.0].domain = self.graph.node(last).domain;
1821 }
1822 for finding in &mut self.graph.numeric_domain_violations {
1824 if finding.library_cite.as_deref() == Some(body.cite) {
1825 finding.edge_or_node = outer;
1826 finding.span = span;
1827 }
1828 }
1829 for finding in &mut self.graph.zero_times_infinity {
1830 if finding.library_cite.as_deref() == Some(body.cite) {
1831 finding.edge_or_node = outer;
1832 finding.span = span;
1833 }
1834 }
1835 self.expansion_cite = None;
1836 self.expanding_library_body = false;
1837 }
1838
1839 fn record_numeric_effects(
1840 &mut self,
1841 id: NodeId,
1842 op: &str,
1843 effect: &OpEffect,
1844 operands: &[NodeId],
1845 span: SrcSpan,
1846 ) {
1847 let library_cite = self.expansion_cite.map(str::to_string);
1848 if let Some(operand) = required_operand(op, effect.requires, operands) {
1849 let producer_domain = self.graph.node(operand).domain;
1850 if let Some(confidence) = domain_violation(effect.requires, producer_domain) {
1851 let proven = matches!(confidence, DomainViolationConfidence::Proven);
1852 let mut message = format!(
1853 "`{}` requires {:?} operand, but producer domain is {:?}{}",
1854 effect.name,
1855 effect.requires,
1856 producer_domain,
1857 if proven {
1858 " with no discharging guard"
1859 } else {
1860 " (producer domain unknown)"
1861 }
1862 );
1863 if let Some(cite) = &library_cite {
1864 message.push_str(&format!(" [expanded from {cite}]"));
1865 }
1866 self.graph
1867 .numeric_domain_violations
1868 .push(NumericDomainViolation {
1869 edge_or_node: id,
1870 op: effect.name.clone(),
1871 requires: format!("{:?}", effect.requires),
1872 producer_domain: format!("{producer_domain:?}"),
1873 proven,
1874 impact: NumericImpact::LocalOnly,
1875 span,
1876 message,
1877 library_cite: library_cite.clone(),
1878 });
1879 }
1880 }
1881
1882 let bare = op.rsplit("::").next().unwrap_or(op);
1883 if matches!(bare, "mul" | "broadcast_mul") && operands.len() >= 2 {
1884 if let Some(mut message) =
1885 zero_times_infinity_message(&self.graph, operands[0], operands[1])
1886 {
1887 if let Some(cite) = &library_cite {
1888 message.push_str(&format!(" [expanded from {cite}]"));
1889 }
1890 self.graph.zero_times_infinity.push(ZeroTimesInfinity {
1891 edge_or_node: id,
1892 impact: NumericImpact::LocalOnly,
1893 span,
1894 message,
1895 library_cite,
1896 });
1897 }
1898 }
1899 }
1900
1901 fn result_domain(&self, effect: &OpEffect, op: &str, operands: &[NodeId]) -> NumericDomain {
1902 let bare = op.rsplit("::").next().unwrap_or(op);
1903 if bare == "affine" {
1904 let operand = operands
1905 .first()
1906 .map(|id| self.graph.node(*id).domain)
1907 .unwrap_or(NumericDomain::Unknown);
1908 let mul = operands
1910 .get(1)
1911 .and_then(|id| literal_f64(&self.graph.node(*id).kind));
1912 let add = operands
1913 .get(2)
1914 .and_then(|id| literal_f64(&self.graph.node(*id).kind));
1915 return affine_domain(operand, mul, add);
1916 }
1917 match bare {
1918 "mul" | "broadcast_mul" | "add" | "broadcast_add" if operands.len() >= 2 => {
1919 op_semantics::join_domain(
1920 self.graph.node(operands[0]).domain,
1921 self.graph.node(operands[1]).domain,
1922 )
1923 }
1924 _ => effect.domain,
1925 }
1926 }
1927}
1928
1929fn required_operand(op: &str, requires: DomainRequirement, operands: &[NodeId]) -> Option<NodeId> {
1930 if matches!(requires, DomainRequirement::None) || operands.is_empty() {
1931 return None;
1932 }
1933 let bare = op.rsplit("::").next().unwrap_or(op);
1934 if matches!(bare, "div" | "broadcast_div") && operands.len() >= 2 {
1936 return Some(operands[1]);
1937 }
1938 Some(operands[0])
1939}
1940
1941fn literal_f64(kind: &NodeKind) -> Option<f64> {
1942 let NodeKind::Literal { text } = kind else {
1943 return None;
1944 };
1945 let trimmed = text.trim().trim_end_matches('_');
1946 let trimmed = trimmed
1947 .strip_suffix("f64")
1948 .or_else(|| trimmed.strip_suffix("f32"))
1949 .or_else(|| trimmed.strip_suffix("f16"))
1950 .unwrap_or(trimmed);
1951 trimmed.parse::<f64>().ok()
1952}
1953
1954fn is_undischarged_log(graph: &ExprGraph, node: NodeId) -> bool {
1955 let node = graph.node(node);
1956 let NodeKind::Call { callee } = &node.kind else {
1957 return false;
1958 };
1959 let bare = callee.rsplit("::").next().unwrap_or(callee.as_str());
1960 if bare != "log" {
1961 return false;
1962 }
1963 let Some(operand) = graph
1964 .edges
1965 .iter()
1966 .find(|edge| edge.to == node.id && edge.label.as_deref() != Some("module"))
1967 .map(|edge| edge.from)
1968 else {
1969 return false;
1970 };
1971 domain_violation(
1972 DomainRequirement::StrictlyPositive,
1973 graph.node(operand).domain,
1974 )
1975 .is_some()
1976}
1977
1978fn zero_times_infinity_message(graph: &ExprGraph, left: NodeId, right: NodeId) -> Option<String> {
1979 let left_zero = domain_includes_zero(graph.node(left).domain);
1980 let right_zero = domain_includes_zero(graph.node(right).domain);
1981 let left_log = is_undischarged_log(graph, left);
1982 let right_log = is_undischarged_log(graph, right);
1983 if (left_zero && right_log) || (right_zero && left_log) {
1984 Some(
1985 "multiply combines a domain that includes 0 with an undischarged `log`, \
1986 which yields `0 * -inf = NaN` rather than a loud `-inf` loss"
1987 .to_string(),
1988 )
1989 } else {
1990 None
1991 }
1992}
1993
1994fn transfer(
1995 effect: &OpEffect,
1996 operands: &[NodeId],
1997 explicit: Option<AbstractDtype>,
1998 lookup: impl Fn(NodeId) -> (AbstractDtype, GradState),
1999) -> (AbstractDtype, GradState, EdgeKind) {
2000 let dtypes: Vec<AbstractDtype> = operands.iter().map(|id| lookup(*id).0).collect();
2001 let grads: Vec<GradState> = operands.iter().map(|id| lookup(*id).1).collect();
2002
2003 let dtype = match effect.dtype {
2004 DtypeRule::Preserve => dtypes.first().copied().unwrap_or(AbstractDtype::Unknown),
2005 DtypeRule::SameAsInputs => {
2006 let known: Vec<_> = dtypes.iter().copied().filter(|d| d.is_known()).collect();
2007 if known.is_empty() {
2008 AbstractDtype::Unknown
2009 } else if known.iter().all(|d| *d == known[0]) {
2010 known[0]
2011 } else {
2012 AbstractDtype::Unknown
2013 }
2014 }
2015 DtypeRule::Explicit => explicit.unwrap_or(AbstractDtype::Unknown),
2016 DtypeRule::Fixed(dtype) => dtype,
2017 DtypeRule::Unknown => AbstractDtype::Unknown,
2018 };
2019
2020 let (grad, edge_kind) = match effect.grad {
2021 GradFlow::Severs => (GradState::Severed, EdgeKind::Severing),
2022 GradFlow::LayoutDependent => (GradState::LayoutDependent, EdgeKind::Data),
2023 GradFlow::Propagates => (propagate_grad(&grads), EdgeKind::Data),
2024 GradFlow::Unknown => (GradState::Unknown, EdgeKind::Data),
2025 };
2026
2027 let grad = if grads.iter().any(|g| matches!(g, GradState::Severed))
2029 && matches!(effect.grad, GradFlow::Propagates)
2030 {
2031 GradState::Severed
2033 } else {
2034 grad
2035 };
2036
2037 (dtype, grad, edge_kind)
2038}
2039
2040fn propagate_grad(grads: &[GradState]) -> GradState {
2041 if grads.is_empty() {
2042 return GradState::Unknown;
2043 }
2044 if grads.iter().any(|g| matches!(g, GradState::Severed)) {
2045 return GradState::Severed;
2046 }
2047 if grads
2048 .iter()
2049 .any(|g| matches!(g, GradState::LayoutDependent))
2050 {
2051 return GradState::LayoutDependent;
2052 }
2053 if grads
2054 .iter()
2055 .any(|g| matches!(g, GradState::Trainable | GradState::Differentiable))
2056 {
2057 return GradState::Differentiable;
2058 }
2059 if grads.iter().all(|g| matches!(g, GradState::Frozen)) {
2060 return GradState::Frozen;
2061 }
2062 GradState::Unknown
2063}
2064
2065fn first_dtype_mismatch(known: &[AbstractDtype]) -> Option<(AbstractDtype, AbstractDtype)> {
2066 let first = *known.first()?;
2067 known
2068 .iter()
2069 .copied()
2070 .find(|d| *d != first)
2071 .map(|other| (first, other))
2072}
2073
2074fn operand_label(op: &str, index: usize, len: usize) -> String {
2075 if len == 2 {
2076 return if index == 0 {
2077 "lhs".into()
2078 } else {
2079 "rhs".into()
2080 };
2081 }
2082 if index == 0
2083 && !matches!(
2084 op,
2085 "cross_entropy" | "nll" | "mse" | "huber" | "binary_cross_entropy_with_logit"
2086 )
2087 {
2088 return "self".into();
2090 }
2091 format!("arg{index}")
2092}
2093
2094fn join_dtypes(ds: &[AbstractDtype]) -> AbstractDtype {
2095 let known: Vec<_> = ds.iter().copied().filter(|d| d.is_known()).collect();
2096 if known.is_empty() {
2097 AbstractDtype::Unknown
2098 } else if known.iter().all(|d| *d == known[0]) {
2099 known[0]
2100 } else {
2101 AbstractDtype::Unknown
2102 }
2103}
2104
2105fn join_grads(gs: &[GradState]) -> GradState {
2106 if gs.is_empty() {
2107 return GradState::Unknown;
2108 }
2109 let first = gs[0];
2110 if gs.iter().all(|g| *g == first) {
2111 first
2112 } else {
2113 GradState::Unknown
2114 }
2115}
2116
2117fn hint_param_grad(type_text: &str, incoming: GradState) -> GradState {
2118 if !matches!(incoming, GradState::Unknown) {
2119 return incoming;
2120 }
2121 match source_type_base(type_text).as_deref() {
2122 Some("Var" | "candle_core::Var") => GradState::Trainable,
2123 _ => GradState::Unknown,
2124 }
2125}
2126
2127fn source_type_base(text: &str) -> Option<String> {
2128 let ty = syn::parse_str::<syn::Type>(text).ok()?;
2129 innermost_type_name(&ty)
2130}
2131
2132fn resolved_return_type(function: &ImplFn, owner_type: &str) -> Option<String> {
2133 let ty = syn::parse_str::<syn::Type>(&function.return_type).ok()?;
2134 let base = result_inner_base(&ty)?;
2135 if base == "Self" {
2136 Some(owner_type.to_string())
2137 } else {
2138 Some(base)
2139 }
2140}
2141
2142fn result_inner_base(ty: &syn::Type) -> Option<String> {
2143 match ty {
2144 syn::Type::Reference(reference) => result_inner_base(&reference.elem),
2145 syn::Type::Paren(paren) => result_inner_base(&paren.elem),
2146 syn::Type::Group(group) => result_inner_base(&group.elem),
2147 syn::Type::Path(path) => {
2148 let segment = path.path.segments.last()?;
2149 if matches!(
2150 segment.ident.to_string().as_str(),
2151 "Result" | "Option" | "Box" | "Arc"
2152 ) {
2153 let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments else {
2154 return None;
2155 };
2156 return arguments.args.iter().find_map(|argument| match argument {
2157 syn::GenericArgument::Type(inner) => result_inner_base(inner),
2158 _ => None,
2159 });
2160 }
2161 Some(
2162 path.path
2163 .segments
2164 .iter()
2165 .map(|segment| segment.ident.to_string())
2166 .collect::<Vec<_>>()
2167 .join("::"),
2168 )
2169 }
2170 _ => None,
2171 }
2172}
2173
2174fn is_candle_tensor_receiver(type_name: &str) -> bool {
2175 matches!(
2176 type_name.trim_start_matches('&').trim().rsplit("::").next(),
2177 Some("Tensor" | "Var")
2178 )
2179}
2180
2181fn is_non_tensor_tensor_method(method: &str) -> bool {
2182 matches!(
2183 method.rsplit("::").next().unwrap_or(method),
2184 "device"
2185 | "dtype"
2186 | "layout"
2187 | "rank"
2188 | "dims"
2189 | "dim"
2190 | "dims1"
2191 | "dims2"
2192 | "dims3"
2193 | "dims4"
2194 | "dims5"
2195 | "elem_count"
2196 | "stride"
2197 | "is_contiguous"
2198 | "to_scalar"
2199 | "to_vec0"
2200 | "to_vec1"
2201 | "to_vec2"
2202 | "to_vec3"
2203 )
2204}
2205
2206fn normalize_qualified_segments(segments: &[String]) -> String {
2207 segments
2208 .iter()
2209 .map(String::as_str)
2210 .skip_while(|segment| matches!(*segment, "crate" | "self"))
2211 .collect::<Vec<_>>()
2212 .join("::")
2213}
2214
2215fn is_explicit_candle_path(segments: &[String]) -> bool {
2216 matches!(
2217 segments.first().map(String::as_str),
2218 Some(
2219 "candle" | "candle_core" | "candle_nn" | "candle_transformers" | "nn" | "ops" | "loss"
2220 )
2221 )
2222}
2223
2224fn is_scalar_literal(expr: &syn::Expr) -> bool {
2225 match expr {
2226 syn::Expr::Lit(literal) => matches!(
2227 literal.lit,
2228 syn::Lit::Int(_)
2229 | syn::Lit::Float(_)
2230 | syn::Lit::Bool(_)
2231 | syn::Lit::Byte(_)
2232 | syn::Lit::Char(_)
2233 ),
2234 syn::Expr::Unary(unary) => is_scalar_literal(&unary.expr),
2235 syn::Expr::Paren(paren) => is_scalar_literal(&paren.expr),
2236 syn::Expr::Group(group) => is_scalar_literal(&group.expr),
2237 _ => false,
2238 }
2239}
2240
2241fn innermost_type_name(ty: &syn::Type) -> Option<String> {
2242 match ty {
2243 syn::Type::Reference(reference) => innermost_type_name(&reference.elem),
2244 syn::Type::Paren(paren) => innermost_type_name(&paren.elem),
2245 syn::Type::Group(group) => innermost_type_name(&group.elem),
2246 syn::Type::Path(path) => {
2247 let segment = path.path.segments.last()?;
2248 if let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments {
2249 if let Some(inner) = arguments.args.iter().find_map(|argument| match argument {
2250 syn::GenericArgument::Type(inner) => Some(inner),
2251 _ => None,
2252 }) {
2253 return innermost_type_name(inner);
2254 }
2255 }
2256 Some(
2257 path.path
2258 .segments
2259 .iter()
2260 .map(|segment| segment.ident.to_string())
2261 .collect::<Vec<_>>()
2262 .join("::"),
2263 )
2264 }
2265 syn::Type::Tuple(tuple) if tuple.elems.len() == 1 => {
2266 tuple.elems.first().and_then(innermost_type_name)
2267 }
2268 _ => None,
2269 }
2270}
2271
2272fn resolve_entrypoint<'a>(
2273 krate: &'a Crate,
2274 entrypoint: &str,
2275) -> anyhow::Result<(&'a ImplFn, String)> {
2276 let function_candidates = krate.function_candidates(entrypoint);
2277 match function_candidates.as_slice() {
2278 [func] => return Ok((func, String::new())),
2279 [] => {}
2280 _ => {
2281 return Err(anyhow::anyhow!(
2282 "entrypoint `{entrypoint}` is ambiguous ({} free functions)",
2283 function_candidates.len()
2284 ))
2285 }
2286 }
2287 if let Some((ty, method)) = entrypoint.rsplit_once("::") {
2288 let candidates = krate.method_candidates(ty, method);
2289 return match candidates.as_slice() {
2290 [func] => Ok((func, ty.to_string())),
2291 [] => Err(anyhow::anyhow!("entrypoint `{entrypoint}` not found")),
2292 _ => Err(anyhow::anyhow!(
2293 "entrypoint `{entrypoint}` is ambiguous ({} methods); select active Cargo cfg",
2294 candidates.len()
2295 )),
2296 };
2297 }
2298 let hits: Vec<_> = krate
2300 .all_methods()
2301 .filter(|func| func.fn_name == entrypoint)
2302 .collect();
2303 match hits.len() {
2304 1 => {
2305 let func = hits[0];
2306 Ok((func, func.qualified_type_name.clone()))
2307 }
2308 0 => Err(anyhow::anyhow!("entrypoint `{entrypoint}` not found")),
2309 _ => Err(anyhow::anyhow!(
2310 "entrypoint `{entrypoint}` is ambiguous; use Type::method"
2311 )),
2312 }
2313}
2314
2315fn bind_pat(env: &mut HashMap<String, NodeId>, pat: &syn::Pat, value: NodeId) {
2316 match pat {
2317 syn::Pat::Ident(id) => {
2318 env.insert(id.ident.to_string(), value);
2319 }
2320 syn::Pat::Type(t) => bind_pat(env, &t.pat, value),
2321 syn::Pat::Reference(r) => bind_pat(env, &r.pat, value),
2322 syn::Pat::Tuple(t) => {
2323 for p in &t.elems {
2326 bind_pat(env, p, value);
2327 }
2328 }
2329 syn::Pat::TupleStruct(t) => {
2330 for p in &t.elems {
2331 bind_pat(env, p, value);
2332 }
2333 }
2334 syn::Pat::Slice(s) => {
2335 for p in &s.elems {
2336 bind_pat(env, p, value);
2337 }
2338 }
2339 syn::Pat::Wild(_) => {}
2340 _ => {}
2341 }
2342}
2343
2344fn path_ident(p: &syn::ExprPath) -> Option<String> {
2345 if p.path.segments.len() == 1 {
2346 Some(p.path.segments[0].ident.to_string())
2347 } else {
2348 None
2349 }
2350}
2351
2352fn path_last(path: &syn::Path) -> String {
2353 path.segments
2354 .last()
2355 .map(|s| s.ident.to_string())
2356 .unwrap_or_default()
2357}
2358
2359fn path_text(path: &syn::Path) -> String {
2360 path.segments
2361 .iter()
2362 .map(|s| s.ident.to_string())
2363 .collect::<Vec<_>>()
2364 .join("::")
2365}
2366
2367fn lit_text(lit: &syn::Lit) -> String {
2368 match lit {
2369 syn::Lit::Str(s) => s.value(),
2370 syn::Lit::Int(i) => i.to_string(),
2371 syn::Lit::Float(f) => f.to_string(),
2372 syn::Lit::Bool(b) => b.value.to_string(),
2373 _ => load::type_text(lit),
2374 }
2375}
2376
2377fn span_of(file: usize, span: proc_macro2::Span) -> SrcSpan {
2378 load::span_of(file, span)
2379}
2380
2381fn span_binop(op: &syn::BinOp) -> proc_macro2::Span {
2382 op.span()
2383}
2384
2385fn expr_kind_name(expr: &syn::Expr) -> &'static str {
2386 match expr {
2387 syn::Expr::Array(_) => "array",
2388 syn::Expr::Async(_) => "async",
2389 syn::Expr::Cast(_) => "cast",
2390 syn::Expr::Index(_) => "index",
2391 syn::Expr::Range(_) => "range",
2392 syn::Expr::Struct(_) => "struct",
2393 syn::Expr::Repeat(_) => "repeat",
2394 syn::Expr::Unsafe(_) => "unsafe",
2395 syn::Expr::While(_) => "while",
2396 syn::Expr::Loop(_) => "loop",
2397 syn::Expr::ForLoop(_) => "for",
2398 syn::Expr::Break(_) => "break",
2399 syn::Expr::Continue(_) => "continue",
2400 syn::Expr::Yield(_) => "yield",
2401 _ => "expr",
2402 }
2403}