Skip to main content

candle_graph/
dataflow.rs

1//! Expression-level dataflow, dtype, and gradient-connectivity analysis.
2//!
3//! Builds a serializable graph over ordinary crate-local free functions and inherent methods
4//! starting from a chosen entrypoint. Transfer rules live in [`crate::op_semantics`]; anything
5//! not covered is left [`GradState::Unknown`] / [`AbstractDtype::Unknown`] rather than guessed.
6
7use 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/// Gradient / trainability state attached to an expression node.
31#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize)]
32#[serde(rename_all = "snake_case")]
33pub enum GradState {
34    /// A leaf that is expected to receive gradients (e.g. `Var` / trainable parameter).
35    Trainable,
36    /// A leaf that must not receive gradients (frozen weight, constant).
37    Frozen,
38    /// An intermediate on a differentiable path.
39    Differentiable,
40    /// Gradient flow was cut (`detach`, known `apply_op*_no_bwd`, …).
41    Severed,
42    /// Gradient exists only under layout assumptions we cannot prove.
43    LayoutDependent,
44    /// Not enough information.
45    Unknown,
46}
47
48/// What an expression node represents in source.
49#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
50#[serde(tag = "kind", rename_all = "snake_case")]
51pub enum NodeKind {
52    /// Function / method parameter.
53    Param { name: String },
54    /// `let` binding.
55    Local { name: String },
56    /// Result of a call or method call.
57    Call { callee: String },
58    /// Literal / constructor-ish value (dtype may still be known from `DType::…`).
59    Literal { text: String },
60    /// Branch join (phi) of mutually exclusive arms.
61    Phi,
62    /// Return value of the analyzed entry / callee.
63    Return,
64    /// Placeholder when the expression could not be modeled.
65    Unknown { reason: String },
66}
67
68/// Edge role in the expression graph.
69#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
70#[serde(rename_all = "snake_case")]
71pub enum EdgeKind {
72    /// Ordinary data dependence (operand → result).
73    Data,
74    /// Data dependence that severs autograd.
75    Severing,
76    /// Control / branch contribution into a phi.
77    Control,
78}
79
80#[derive(Debug, Clone, Serialize)]
81pub struct ExprNode {
82    pub id: NodeId,
83    pub kind: NodeKind,
84    pub span: SrcSpan,
85    /// Symbolic shape text when observed; never invented.
86    pub shape: Option<String>,
87    pub dtype: AbstractDtype,
88    pub grad: GradState,
89    /// Float-range domain after rounding. Unknown until a catalog rule assigns one.
90    pub domain: NumericDomain,
91    /// Best-effort resolved source type. `None` remains explicitly unknown.
92    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    /// Operand slot or note, e.g. `lhs`, `rhs`, `self`.
102    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/// A same-dtype op where at least one operand is known and another remains unknown.
116///
117/// This is intentionally separate from [`DtypeConflict`]: it highlights a runtime mismatch risk
118/// without claiming that the unknown operand definitely differs.
119#[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/// How a numeric hazard can interfere with training or inference.
129#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Default)]
130#[serde(rename_all = "snake_case")]
131pub enum NumericImpact {
132    /// Hazard is on a loss sink (or reaches one): training can abort with NaN loss.
133    TrainingLossNaN,
134    /// Hazard lies on a trainable → loss path: gradients can be poisoned.
135    GradientPoison,
136    /// Hazard reaches the entry return of a non-loss forward.
137    InferenceOutputRisk,
138    /// Unstable value proven locally but not shown to reach loss or return.
139    #[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/// A partial function applied to a domain that can attain a forbidden endpoint.
161#[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    /// Source citation when the hazard came from an expanded library body.
172    #[serde(skip_serializing_if = "Option::is_none")]
173    pub library_cite: Option<String>,
174}
175
176/// Multiply that can evaluate `0 * ±inf` (silent NaN) rather than a loud `-inf` loss.
177#[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    /// Entrypoint return node, when analysis produced one.
198    pub entry_return: Option<NodeId>,
199    /// Nodes that are loss sinks (`cross_entropy`, `mse`, …).
200    pub loss_nodes: Vec<NodeId>,
201    /// Trainable / frozen parameter leaves discovered during analysis.
202    pub param_nodes: Vec<NodeId>,
203    /// Nodes proven to be Candle tensor values. Non-tensor expressions stay out of ModelIr.
204    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    /// Trainable leaves with no reverse data path to any loss node.
223    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    /// Edges marked as severing autograd.
235    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    /// Same-dtype operand mismatches recorded during transfer.
244    pub fn dtype_conflicts(&self) -> &[DtypeConflict] {
245        &self.dtype_conflicts
246    }
247
248    /// Partially-known operands at an op that requires matching dtypes.
249    pub fn dtype_risks(&self) -> &[DtypeRisk] {
250        &self.dtype_risks
251    }
252
253    /// Differentiable paths from `from` to `to`; severing edges are intentionally excluded.
254    /// Caps path count to keep recursion honest on dense graphs.
255    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    /// Nodes that can reach a loss following edges forward (operand → result → … → loss).
316    fn nodes_reaching_losses(&self) -> HashSet<NodeId> {
317        // Reverse adjacency: result → operands, then BFS from losses.
318        let mut rev: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
319        for e in &self.edges {
320            if matches!(e.kind, EdgeKind::Severing) {
321                // Severed edges do not carry gradient; they do not count for "alive".
322                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    /// Forward-reachable nodes from `start` along non-severing edges.
346    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    /// Classify numeric hazards by whether they can fail training or interfere with inference.
371    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
462/// Analyze `entrypoint` against a loaded crate.
463///
464/// `entrypoint` is either a free function name (`forward`) or an inherent method
465/// (`Model::forward`).
466pub fn analyze(krate: &Crate, entrypoint: &str) -> anyhow::Result<ExprGraph> {
467    analyze_with_candle_version(krate, entrypoint, None)
468}
469
470/// Analyze with the resolved candle-nn version used to gate version-specific gradient rules.
471pub 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    /// Best-effort source types for expression nodes. Entries are only added from explicit
486    /// signatures, struct fields, or known Candle tensor operations; absence means unknown.
487    node_types: HashMap<NodeId, String>,
488    /// Recursion guard: `(type_name, fn_name)` currently on the call stack.
489    call_stack: HashSet<(String, String)>,
490    file_hint: usize,
491    module_path: String,
492    candle_nn_version: Option<String>,
493    /// True while expanding an audited library body (prevents recursive expansion).
494    expanding_library_body: bool,
495    /// Citation attached to numeric findings produced by the active expansion.
496    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                // Only an explicit `Var` parameter type proves trainability. Tensor names such as
678                // `weight` do not: frozen and trainable tensors share the same type.
679                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                // Field projection: dtype/grad Unknown unless we know more — honest default.
795                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                // Do not infer trainability from a struct field called `weight`: candle models
816                // commonly store frozen base weights and trainable adapters side by side. Without
817                // a resolved builder root, Unknown is safer than a false dead-parameter report.
818                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                // Analyze one symbolic iteration. This preserves the body dataflow without
827                // pretending to know the runtime trip count.
828                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                // Analyze closure body in a forked env; captures stay as outer bindings.
895                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                // Recognize DType::… inside simple macros? Leave Unknown.
932                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        // Only an actual `DType::<variant>` path is dtype evidence. A constant or user type whose
995        // name happens to end in `F32` must not affect tensor inference.
996        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        // Write back assignments to names that existed in the outer env.
1029        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        // Inherent crate-local method on a known type? Only when receiver looks like `self`
1225        // or we can recover a type name — keep interprocedural for `Self::` style via func_call.
1226        // For `self.foo(...)` try type of `self` from env naming: look up methods by scanning.
1227        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        // Link only the callee result. Direct actual-argument edges would bypass detach/no-bwd
1373        // operations inside the callee and create false live-gradient paths.
1374        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        // Path call: Type::method or free function or candle_nn::loss::cross_entropy
1409        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            // Inherent Type::method
1422            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            // Free function
1442            let function_name = normalize_qualified_segments(&segs);
1443            let candidates = self.krate.function_candidates(&function_name);
1444            if let [func] = candidates.as_slice() {
1445                // Prefer candle_nn loss / op names when path mentions candle_nn — still apply
1446                // transfer rules on the known last segment either way.
1447                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                    // If it's a known op AND a local function, the local body wins for
1459                    // interprocedural detail; still record the op effect on the call node.
1460                }
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            // Known library op (candle_nn::loss::cross_entropy, etc.)
1473            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            // Unknown external call
1488            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        // Also apply transfer rule when the function name itself is a known op (e.g. a
1551        // thin local wrapper is uncommon; known candle_nn names take the op path above).
1552        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    /// Apply a named op transfer rule to `operands` (receiver first for methods).
1588    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, // honest: conflicting inputs
1630                    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    /// Expand an audited library body into synthetic ops judged by the same domain pass.
1745    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        // Point findings at the user call site while keeping synthetic nodes for domain facts.
1823        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            // Method form: receiver, mul, add — epsilon guard needs mul > 0 and add > 0.
1909            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    // For division the non-zero requirement applies to the divisor.
1935    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    // Once severed, stay severed even if rule said propagates — handled by edge kind on inputs.
2028    let grad = if grads.iter().any(|g| matches!(g, GradState::Severed))
2029        && matches!(effect.grad, GradFlow::Propagates)
2030    {
2031        // Inputs already severed: result is severed for connectivity purposes.
2032        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        // method receiver
2089        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    // Bare method name unique across all loaded impls?
2299    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            // Without splitting the value node, bind each subpat to the same node (honest Unknown
2324            // would be worse for simple `(a, b) = …` where we lack tuple projection).
2325            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}