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