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