Skip to main content

candle_graph/
dataflow.rs

1//! Expression-level dataflow, dtype, and gradient-connectivity analysis.
2//!
3//! Builds a serializable graph over ordinary crate-local free functions and inherent methods
4//! starting from a chosen entrypoint. Transfer rules live in [`crate::op_semantics`]; anything
5//! not covered is left [`GradState::Unknown`] / [`AbstractDtype::Unknown`] rather than guessed.
6
7use std::collections::{HashMap, HashSet, VecDeque};
8
9use serde::Serialize;
10use syn::spanned::Spanned;
11
12use crate::ir::SrcSpan;
13use crate::load::{self, Crate, ImplFn};
14use crate::op_semantics::{
15    self, affine_domain, domain_includes_zero, domain_violation, library_body, AbstractDtype,
16    BodyAtom, DomainRequirement, DomainViolationConfidence, DtypeRule, GradFlow, LibraryBody,
17    NumericDomain, OpEffect,
18};
19use crate::phase::ExecutionPhase;
20
21macro_rules! id_type {
22    ($name:ident) => {
23        #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize)]
24        pub struct $name(pub usize);
25    };
26}
27
28id_type!(NodeId);
29id_type!(EdgeId);
30
31/// 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        if method == "dtype"
1251            && self
1252                .node_type(receiver)
1253                .is_some_and(is_candle_tensor_receiver)
1254        {
1255            let dtype = self.graph.node(receiver).dtype;
1256            return Ok(self.add_node(
1257                NodeKind::Literal {
1258                    text: "dtype()".into(),
1259                },
1260                span,
1261                dtype,
1262                GradState::Frozen,
1263                None,
1264            ));
1265        }
1266
1267        if self
1268            .node_type(receiver)
1269            .is_some_and(|receiver_type| receiver_type.contains("VarBuilder"))
1270        {
1271            if matches!(method.as_str(), "pp" | "push_prefix" | "clone") {
1272                let dtype = self.graph.node(receiver).dtype;
1273                let id = self.add_node(
1274                    NodeKind::Call {
1275                        callee: method.clone(),
1276                    },
1277                    span,
1278                    dtype,
1279                    GradState::Frozen,
1280                    None,
1281                );
1282                self.set_node_type(id, "VarBuilder");
1283                self.add_edge(receiver, id, EdgeKind::Data, Some("self".into()));
1284                for (index, arg) in args.iter().enumerate() {
1285                    self.add_edge(*arg, id, EdgeKind::Data, Some(format!("arg{index}")));
1286                }
1287                return Ok(id);
1288            }
1289            if matches!(
1290                method.as_str(),
1291                "get"
1292                    | "get_with_hints"
1293                    | "get_with_hints_dtype"
1294                    | "get_unchecked"
1295                    | "get_unchecked_dtype"
1296                    | "get_with_dtype"
1297            ) {
1298                let dtype = self.graph.node(receiver).dtype;
1299                if dtype.is_known() {
1300                    let id = self.add_node(
1301                        NodeKind::Call {
1302                            callee: method.clone(),
1303                        },
1304                        span,
1305                        dtype,
1306                        GradState::Trainable,
1307                        None,
1308                    );
1309                    self.set_node_type(id, "Tensor");
1310                    self.add_edge(receiver, id, EdgeKind::Data, Some("builder".into()));
1311                    for (index, arg) in args.iter().enumerate() {
1312                        self.add_edge(*arg, id, EdgeKind::Data, Some(format!("arg{index}")));
1313                    }
1314                    return Ok(id);
1315                }
1316            }
1317        }
1318
1319        // Inherent crate-local method on a known type? Only when receiver looks like `self`
1320        // or we can recover a type name — keep interprocedural for `Self::` style via func_call.
1321        // For `self.foo(...)` try type of `self` from env naming: look up methods by scanning.
1322        if let Some(ret) = self.try_interproc_method(&method, receiver, &args, span)? {
1323            return Ok(ret);
1324        }
1325
1326        if let Some(receiver_type) = self.node_type(receiver).map(str::to_string) {
1327            let effect = op_semantics::lookup_method(
1328                &receiver_type,
1329                &method,
1330                self.candle_nn_version.as_deref(),
1331            );
1332            if !matches!(effect.dtype, DtypeRule::Unknown)
1333                || !matches!(effect.grad, GradFlow::Unknown)
1334            {
1335                let label = effect.name.clone();
1336                let result = self.apply_effect(&label, effect, span, &args, None);
1337                if !is_candle_tensor_receiver(&receiver_type) {
1338                    self.add_edge(receiver, result, EdgeKind::Data, Some("module".into()));
1339                }
1340                return Ok(result);
1341            }
1342        }
1343
1344        let mut operands = vec![receiver];
1345        operands.extend(args.iter().copied());
1346        if !self
1347            .node_type(receiver)
1348            .is_some_and(is_candle_tensor_receiver)
1349        {
1350            let id = self.add_node(
1351                NodeKind::Call {
1352                    callee: method.clone(),
1353                },
1354                span,
1355                AbstractDtype::Unknown,
1356                GradState::Unknown,
1357                None,
1358            );
1359            self.add_edge(receiver, id, EdgeKind::Data, Some("self".into()));
1360            for (index, arg) in args.iter().enumerate() {
1361                self.add_edge(*arg, id, EdgeKind::Data, Some(format!("arg{index}")));
1362            }
1363            self.diagnose(
1364                span,
1365                format!(
1366                    "receiver type for method `{method}` is not proven to be Tensor or a \
1367                     crate-local type; transfer semantics left unknown"
1368                ),
1369            );
1370            return Ok(id);
1371        }
1372        let explicit = if method == "to_dtype" {
1373            args.first().map(|id| self.graph.node(*id).dtype)
1374        } else {
1375            None
1376        };
1377        Ok(self.apply_op(&method, span, &operands, explicit))
1378    }
1379
1380    fn try_interproc_method(
1381        &mut self,
1382        method: &str,
1383        receiver: NodeId,
1384        args: &[NodeId],
1385        span: SrcSpan,
1386    ) -> anyhow::Result<Option<NodeId>> {
1387        let Some(receiver_type) = self.node_type(receiver).map(ToString::to_string) else {
1388            let count = self
1389                .krate
1390                .all_methods()
1391                .filter(|func| func.fn_name == method)
1392                .count();
1393            if count > 0 {
1394                self.diagnose(
1395                    span,
1396                    format!(
1397                        "method `{method}` has {count} crate-local candidate(s), but the receiver \
1398                         type is unknown; not inlined"
1399                    ),
1400                );
1401            }
1402            return Ok(None);
1403        };
1404        if is_candle_tensor_receiver(&receiver_type) {
1405            return Ok(None);
1406        }
1407
1408        let candidates = self.krate.method_candidates(&receiver_type, method);
1409        let func = match candidates.as_slice() {
1410            [] => return Ok(None),
1411            [func] => *func,
1412            _ => {
1413                self.diagnose(
1414                    span,
1415                    format!(
1416                        "ambiguous method `{receiver_type}::{method}` ({} candidates); not inlined",
1417                        candidates.len()
1418                    ),
1419                );
1420                return Ok(None);
1421            }
1422        };
1423
1424        let key = (func.type_name.clone(), func.fn_name.clone());
1425        if self.call_stack.contains(&key) {
1426            self.diagnose(
1427                span,
1428                format!("recursion guard: {}::{} already on stack", key.0, key.1),
1429            );
1430            let id = self.add_node(
1431                NodeKind::Call {
1432                    callee: format!("{}::{method}", func.type_name),
1433                },
1434                span,
1435                AbstractDtype::Unknown,
1436                GradState::Unknown,
1437                None,
1438            );
1439            self.add_edge(receiver, id, EdgeKind::Data, Some("self".into()));
1440            for (i, a) in args.iter().enumerate() {
1441                self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1442            }
1443            return Ok(Some(id));
1444        }
1445
1446        self.call_stack.insert(key.clone());
1447        let prev_file = self.file_hint;
1448        let prev_module = self.module_path.clone();
1449        self.file_hint = func.span.file;
1450        self.module_path = func.module_path.clone();
1451        let ret = self.analyze_function(func, &func.type_name, Some(receiver), args)?;
1452        self.file_hint = prev_file;
1453        self.module_path = prev_module;
1454        self.call_stack.remove(&key);
1455
1456        let out = ret.unwrap_or_else(|| {
1457            self.add_node(
1458                NodeKind::Call {
1459                    callee: format!("{}::{method}", func.type_name),
1460                },
1461                span,
1462                AbstractDtype::Unknown,
1463                GradState::Unknown,
1464                None,
1465            )
1466        });
1467        // Link only the callee result. Direct actual-argument edges would bypass detach/no-bwd
1468        // operations inside the callee and create false live-gradient paths.
1469        let call_node = self.add_node(
1470            NodeKind::Call {
1471                callee: format!("{}::{method}", func.type_name),
1472            },
1473            span,
1474            self.graph.node(out).dtype,
1475            self.graph.node(out).grad,
1476            self.graph.node(out).shape.clone(),
1477        );
1478        if let Some(return_type) = resolved_return_type(func, &func.type_name) {
1479            self.set_node_type(call_node, return_type);
1480        }
1481        self.add_edge(out, call_node, EdgeKind::Data, Some("return".into()));
1482        Ok(Some(call_node))
1483    }
1484
1485    fn func_call(
1486        &mut self,
1487        env: &mut HashMap<String, NodeId>,
1488        c: &syn::ExprCall,
1489    ) -> anyhow::Result<NodeId> {
1490        let mut args = Vec::with_capacity(c.args.len());
1491        for a in &c.args {
1492            args.push(self.expr(env, a)?);
1493        }
1494        let span = match &*c.func {
1495            syn::Expr::Path(p) => span_of(self.file_hint, p.path.span()),
1496            _ => SrcSpan {
1497                file: self.file_hint,
1498                line: 0,
1499                col: 0,
1500            },
1501        };
1502
1503        // Path call: Type::method or free function or candle_nn::loss::cross_entropy
1504        if let syn::Expr::Path(p) = &*c.func {
1505            let source_segments: Vec<String> = p
1506                .path
1507                .segments
1508                .iter()
1509                .map(|s| s.ident.to_string())
1510                .collect();
1511            let segs = self
1512                .krate
1513                .resolve_import_path(&self.module_path, &source_segments);
1514            let last = segs.last().cloned().unwrap_or_default();
1515
1516            if matches!(
1517                last.as_str(),
1518                "from_varmap" | "from_mmaped_safetensors" | "from_buffered_safetensors"
1519            ) && segs.iter().any(|segment| segment == "VarBuilder")
1520            {
1521                let dtype = c
1522                    .args
1523                    .get(1)
1524                    .map(|expression| self.expr(env, expression))
1525                    .transpose()?
1526                    .map(|node| self.graph.node(node).dtype)
1527                    .filter(|dtype| dtype.is_known())
1528                    .unwrap_or(AbstractDtype::Unknown);
1529                let id = self.add_node(
1530                    NodeKind::Call {
1531                        callee: segs.join("::"),
1532                    },
1533                    span,
1534                    dtype,
1535                    GradState::Frozen,
1536                    None,
1537                );
1538                self.set_node_type(id, "VarBuilder");
1539                for (index, arg) in args.iter().enumerate() {
1540                    self.add_edge(*arg, id, EdgeKind::Data, Some(format!("arg{index}")));
1541                }
1542                return Ok(id);
1543            }
1544
1545            if is_tensor_constructor_name(&last) && segs.iter().any(|segment| segment == "Tensor") {
1546                return self.tensor_constructor(&last, &args, &c.args, span);
1547            }
1548
1549            // Inherent Type::method
1550            if segs.len() >= 2 {
1551                let type_name = normalize_qualified_segments(&segs[..segs.len() - 1]);
1552                let candidates = self.krate.method_candidates(&type_name, &last);
1553                match candidates.as_slice() {
1554                    [func] => return self.call_crate_fn(func, &type_name, None, &args, span),
1555                    [] => {}
1556                    _ => {
1557                        self.diagnose(
1558                            span,
1559                            format!(
1560                                "call `{type_name}::{last}` is ambiguous ({} definitions); not \
1561                                 inlined",
1562                                candidates.len()
1563                            ),
1564                        );
1565                    }
1566                }
1567            }
1568
1569            // Free function
1570            let function_name = normalize_qualified_segments(&segs);
1571            let candidates = self.krate.function_candidates(&function_name);
1572            if let [func] = candidates.as_slice() {
1573                // Prefer candle_nn loss / op names when path mentions candle_nn — still apply
1574                // transfer rules on the known last segment either way.
1575                if op_semantics::lookup_for(&last, self.candle_nn_version.as_deref())
1576                    .note
1577                    .is_some()
1578                    || matches!(
1579                        op_semantics::lookup_for(&last, self.candle_nn_version.as_deref()).dtype,
1580                        DtypeRule::SameAsInputs
1581                            | DtypeRule::Preserve
1582                            | DtypeRule::Explicit
1583                            | DtypeRule::Fixed(_)
1584                    )
1585                {
1586                    // If it's a known op AND a local function, the local body wins for
1587                    // interprocedural detail; still record the op effect on the call node.
1588                }
1589                return self.call_crate_fn(func, "", None, &args, span);
1590            } else if candidates.len() > 1 {
1591                self.diagnose(
1592                    span,
1593                    format!(
1594                        "free function `{function_name}` is ambiguous ({} definitions); not inlined",
1595                        candidates.len()
1596                    ),
1597                );
1598            }
1599
1600            // Known library op (candle_nn::loss::cross_entropy, etc.)
1601            let effect = op_semantics::lookup_for(&last, self.candle_nn_version.as_deref());
1602            if is_explicit_candle_path(&segs)
1603                && (effect.note.is_some()
1604                    || !matches!(effect.grad, GradFlow::Unknown)
1605                    || !matches!(effect.dtype, DtypeRule::Unknown))
1606            {
1607                let explicit = if last == "to_dtype" {
1608                    args.first().map(|id| self.graph.node(*id).dtype)
1609                } else {
1610                    None
1611                };
1612                return Ok(self.apply_op(&last, span, &args, explicit));
1613            }
1614
1615            // Unknown external call
1616            let id = self.add_node(
1617                NodeKind::Call {
1618                    callee: segs.join("::"),
1619                },
1620                span,
1621                AbstractDtype::Unknown,
1622                GradState::Unknown,
1623                None,
1624            );
1625            for (i, a) in args.iter().enumerate() {
1626                self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1627            }
1628            return Ok(id);
1629        }
1630
1631        let callee = self.expr(env, &c.func)?;
1632        let id = self.add_node(
1633            NodeKind::Call {
1634                callee: "call".into(),
1635            },
1636            span,
1637            AbstractDtype::Unknown,
1638            GradState::Unknown,
1639            None,
1640        );
1641        self.add_edge(callee, id, EdgeKind::Data, Some("callee".into()));
1642        for (i, a) in args.iter().enumerate() {
1643            self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1644        }
1645        Ok(id)
1646    }
1647
1648    fn tensor_constructor(
1649        &mut self,
1650        name: &str,
1651        arg_nodes: &[NodeId],
1652        arg_exprs: &syn::punctuated::Punctuated<syn::Expr, syn::token::Comma>,
1653        span: SrcSpan,
1654    ) -> anyhow::Result<NodeId> {
1655        let dtype = match name {
1656            "zeros" | "ones" | "zeros_like" | "ones_like" => {
1657                let from_syn = arg_exprs
1658                    .iter()
1659                    .nth(if name.ends_with("_like") { 0 } else { 1 })
1660                    .and_then(dtype_from_syn_expr)
1661                    .filter(|dtype| dtype.is_known());
1662                let from_node = if name.ends_with("_like") {
1663                    arg_nodes
1664                        .first()
1665                        .map(|node| self.graph.node(*node).dtype)
1666                        .filter(|dtype| dtype.is_known())
1667                } else {
1668                    arg_nodes
1669                        .get(1)
1670                        .map(|node| self.graph.node(*node).dtype)
1671                        .filter(|dtype| dtype.is_known())
1672                };
1673                from_node.or(from_syn).unwrap_or(AbstractDtype::Unknown)
1674            }
1675            "from_vec" | "new" => arg_exprs
1676                .first()
1677                .and_then(dtype_from_collection_expr)
1678                .or_else(|| {
1679                    arg_nodes
1680                        .first()
1681                        .map(|node| self.graph.node(*node).dtype)
1682                })
1683                .filter(|dtype| dtype.is_known())
1684                .unwrap_or(AbstractDtype::Unknown),
1685            "arange" | "arange_step" => arg_exprs
1686                .first()
1687                .and_then(dtype_from_scalar_expr)
1688                .filter(|dtype| dtype.is_known())
1689                .unwrap_or(AbstractDtype::Unknown),
1690            "rand" | "randn" => arg_exprs
1691                .first()
1692                .and_then(dtype_from_scalar_expr)
1693                .filter(|dtype| dtype.is_known())
1694                .unwrap_or(AbstractDtype::Unknown),
1695            _ => AbstractDtype::Unknown,
1696        };
1697        let id = self.add_node(
1698            NodeKind::Call {
1699                callee: format!("Tensor::{name}"),
1700            },
1701            span,
1702            dtype,
1703            GradState::Frozen,
1704            None,
1705        );
1706        self.set_node_type(id, "Tensor");
1707        for (index, arg) in arg_nodes.iter().enumerate() {
1708            self.add_edge(*arg, id, EdgeKind::Data, Some(format!("arg{index}")));
1709        }
1710        Ok(id)
1711    }
1712
1713    fn call_crate_fn(
1714        &mut self,
1715        func: &ImplFn,
1716        type_name: &str,
1717        receiver: Option<NodeId>,
1718        args: &[NodeId],
1719        span: SrcSpan,
1720    ) -> anyhow::Result<NodeId> {
1721        let key = (func.type_name.clone(), func.fn_name.clone());
1722        let label = if type_name.is_empty() {
1723            func.fn_name.clone()
1724        } else {
1725            format!("{type_name}::{}", func.fn_name)
1726        };
1727
1728        if self.call_stack.contains(&key) {
1729            self.diagnose(span, format!("recursion guard: {label} already on stack"));
1730            let id = self.add_node(
1731                NodeKind::Call { callee: label },
1732                span,
1733                AbstractDtype::Unknown,
1734                GradState::Unknown,
1735                None,
1736            );
1737            for (i, a) in args.iter().enumerate() {
1738                self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1739            }
1740            return Ok(id);
1741        }
1742
1743        // Also apply transfer rule when the function name itself is a known op (e.g. a
1744        // thin local wrapper is uncommon; known candle_nn names take the op path above).
1745        self.call_stack.insert(key.clone());
1746        let prev_file = self.file_hint;
1747        let prev_module = self.module_path.clone();
1748        self.file_hint = func.span.file;
1749        self.module_path = func.module_path.clone();
1750        let ret = self.analyze_function(func, &func.type_name, receiver, args)?;
1751        self.file_hint = prev_file;
1752        self.module_path = prev_module;
1753        self.call_stack.remove(&key);
1754
1755        let out = ret.unwrap_or_else(|| {
1756            self.add_node(
1757                NodeKind::Call {
1758                    callee: label.clone(),
1759                },
1760                span,
1761                AbstractDtype::Unknown,
1762                GradState::Unknown,
1763                None,
1764            )
1765        });
1766        let call_node = self.add_node(
1767            NodeKind::Call { callee: label },
1768            span,
1769            self.graph.node(out).dtype,
1770            self.graph.node(out).grad,
1771            self.graph.node(out).shape.clone(),
1772        );
1773        if let Some(return_type) = resolved_return_type(func, type_name) {
1774            self.set_node_type(call_node, return_type);
1775        }
1776        self.add_edge(out, call_node, EdgeKind::Data, Some("return".into()));
1777        Ok(call_node)
1778    }
1779
1780    /// Apply a named op transfer rule to `operands` (receiver first for methods).
1781    fn apply_op(
1782        &mut self,
1783        op: &str,
1784        span: SrcSpan,
1785        operands: &[NodeId],
1786        explicit_dtype: Option<AbstractDtype>,
1787    ) -> NodeId {
1788        let effect = op_semantics::lookup_for(op, self.candle_nn_version.as_deref());
1789        self.apply_effect(op, effect, span, operands, explicit_dtype)
1790    }
1791
1792    fn apply_effect(
1793        &mut self,
1794        op: &str,
1795        effect: OpEffect,
1796        span: SrcSpan,
1797        operands: &[NodeId],
1798        explicit_dtype: Option<AbstractDtype>,
1799    ) -> NodeId {
1800        let (dtype, grad, edge_kind) = transfer(&effect, operands, explicit_dtype, |id| {
1801            let n = &self.graph.nodes[id.0];
1802            (n.dtype, n.grad)
1803        });
1804        let domain = self.result_domain(&effect, op, operands);
1805
1806        if matches!(effect.dtype, DtypeRule::SameAsInputs) {
1807            let operand_dtypes: Vec<AbstractDtype> = operands
1808                .iter()
1809                .map(|id| self.graph.node(*id).dtype)
1810                .collect();
1811            let known: Vec<AbstractDtype> = operand_dtypes
1812                .iter()
1813                .copied()
1814                .filter(|d| d.is_known())
1815                .collect();
1816            if let Some((a, b)) = first_dtype_mismatch(&known) {
1817                let id = self.add_node_with_domain(
1818                    NodeKind::Call {
1819                        callee: op.to_string(),
1820                    },
1821                    span,
1822                    AbstractDtype::Unknown, // honest: conflicting inputs
1823                    grad,
1824                    None,
1825                    domain,
1826                );
1827                self.graph.dtype_conflicts.push(DtypeConflict {
1828                    edge_or_node: id,
1829                    op: op.to_string(),
1830                    left: a,
1831                    right: b,
1832                    span,
1833                    message: format!("{op} requires same dtype, got {a} vs {b}"),
1834                });
1835                for (i, &src) in operands.iter().enumerate() {
1836                    let label = operand_label(op, i, operands.len());
1837                    self.add_edge(src, id, edge_kind, Some(label));
1838                }
1839                if effect.is_loss {
1840                    self.graph.loss_nodes.push(id);
1841                }
1842                self.set_node_type(id, "Tensor");
1843                self.finish_numeric_call(id, op, &effect, operands, span);
1844                return id;
1845            }
1846            if let Some(known_dtype) = known.first().copied().filter(|_| {
1847                operand_dtypes
1848                    .iter()
1849                    .any(|dtype| !matches!(dtype, AbstractDtype::Unknown))
1850                    && operand_dtypes
1851                        .iter()
1852                        .any(|dtype| matches!(dtype, AbstractDtype::Unknown))
1853            }) {
1854                let id = self.add_node_with_domain(
1855                    NodeKind::Call {
1856                        callee: op.to_string(),
1857                    },
1858                    span,
1859                    dtype,
1860                    grad,
1861                    None,
1862                    domain,
1863                );
1864                self.graph.dtype_risks.push(DtypeRisk {
1865                    edge_or_node: id,
1866                    op: op.to_string(),
1867                    known: known_dtype,
1868                    span,
1869                    message: format!(
1870                        "{op} requires matching dtypes; one operand is {known_dtype} and another is unknown"
1871                    ),
1872                });
1873                for (i, &src) in operands.iter().enumerate() {
1874                    let label = operand_label(op, i, operands.len());
1875                    self.add_edge(src, id, edge_kind, Some(label));
1876                }
1877                if effect.is_loss {
1878                    self.graph.loss_nodes.push(id);
1879                }
1880                self.set_node_type(id, "Tensor");
1881                self.finish_numeric_call(id, op, &effect, operands, span);
1882                return id;
1883            }
1884        }
1885
1886        let id = self.add_node_with_domain(
1887            NodeKind::Call {
1888                callee: op.to_string(),
1889            },
1890            span,
1891            dtype,
1892            grad,
1893            None,
1894            domain,
1895        );
1896        for (i, &src) in operands.iter().enumerate() {
1897            let label = operand_label(op, i, operands.len());
1898            self.add_edge(src, id, edge_kind, Some(label));
1899        }
1900        if effect.is_loss {
1901            self.graph.loss_nodes.push(id);
1902        }
1903        if !matches!(effect.dtype, DtypeRule::Unknown)
1904            || !matches!(effect.grad, GradFlow::Unknown)
1905            || operands.iter().any(|operand| {
1906                self.node_type(*operand)
1907                    .is_some_and(is_candle_tensor_receiver)
1908            }) && !is_non_tensor_tensor_method(op)
1909        {
1910            self.set_node_type(id, "Tensor");
1911        }
1912        if let Some(note) = effect.note {
1913            if matches!(effect.grad, GradFlow::Unknown) {
1914                self.diagnose(span, note);
1915            }
1916        }
1917        self.finish_numeric_call(id, op, &effect, operands, span);
1918        id
1919    }
1920
1921    fn finish_numeric_call(
1922        &mut self,
1923        id: NodeId,
1924        op: &str,
1925        effect: &OpEffect,
1926        operands: &[NodeId],
1927        span: SrcSpan,
1928    ) {
1929        self.record_numeric_effects(id, op, effect, operands, span);
1930        if !self.expanding_library_body {
1931            if let Some(body) = library_body(op, self.candle_nn_version.as_deref()) {
1932                self.expand_library_body(id, body, operands, span);
1933            }
1934        }
1935    }
1936
1937    /// Expand an audited library body into synthetic ops judged by the same domain pass.
1938    fn expand_library_body(
1939        &mut self,
1940        outer: NodeId,
1941        body: &LibraryBody,
1942        args: &[NodeId],
1943        span: SrcSpan,
1944    ) {
1945        self.expanding_library_body = true;
1946        self.expansion_cite = Some(body.cite);
1947        let mut vals: Vec<NodeId> = Vec::with_capacity(body.steps.len());
1948        for step in body.steps {
1949            let id = match *step {
1950                BodyAtom::Arg(index) => args.get(index).copied().unwrap_or_else(|| {
1951                    self.add_node(
1952                        NodeKind::Unknown {
1953                            reason: format!("missing library body arg {index}"),
1954                        },
1955                        span,
1956                        AbstractDtype::Unknown,
1957                        GradState::Unknown,
1958                        None,
1959                    )
1960                }),
1961                BodyAtom::Assume { src, domain } => {
1962                    let src = vals[src as usize];
1963                    let id = self.add_node_with_domain(
1964                        NodeKind::Call {
1965                            callee: "library_domain_assume".into(),
1966                        },
1967                        span,
1968                        self.graph.node(src).dtype,
1969                        self.graph.node(src).grad,
1970                        None,
1971                        domain,
1972                    );
1973                    self.add_edge(src, id, EdgeKind::Data, Some("assume".into()));
1974                    self.set_node_type(id, "Tensor");
1975                    id
1976                }
1977                BodyAtom::Unary { op, src } => {
1978                    let src = vals[src as usize];
1979                    self.apply_op(op, span, &[src], None)
1980                }
1981                BodyAtom::Binary { op, left, right } => {
1982                    let left = vals[left as usize];
1983                    let right = vals[right as usize];
1984                    self.apply_op(op, span, &[left, right], None)
1985                }
1986                BodyAtom::Affine { src, mul, add } => {
1987                    let src = vals[src as usize];
1988                    let mul_node = self.add_node(
1989                        NodeKind::Literal {
1990                            text: format!("{mul}"),
1991                        },
1992                        span,
1993                        AbstractDtype::Unknown,
1994                        GradState::Frozen,
1995                        None,
1996                    );
1997                    let add_node = self.add_node(
1998                        NodeKind::Literal {
1999                            text: format!("{add}"),
2000                        },
2001                        span,
2002                        AbstractDtype::Unknown,
2003                        GradState::Frozen,
2004                        None,
2005                    );
2006                    self.apply_op("affine", span, &[src, mul_node, add_node], None)
2007                }
2008            };
2009            vals.push(id);
2010        }
2011        if let Some(&last) = vals.last() {
2012            self.add_edge(last, outer, EdgeKind::Data, Some("expanded_body".into()));
2013            self.graph.nodes[outer.0].domain = self.graph.node(last).domain;
2014        }
2015        // Point findings at the user call site while keeping synthetic nodes for domain facts.
2016        for finding in &mut self.graph.numeric_domain_violations {
2017            if finding.library_cite.as_deref() == Some(body.cite) {
2018                finding.edge_or_node = outer;
2019                finding.span = span;
2020            }
2021        }
2022        for finding in &mut self.graph.zero_times_infinity {
2023            if finding.library_cite.as_deref() == Some(body.cite) {
2024                finding.edge_or_node = outer;
2025                finding.span = span;
2026            }
2027        }
2028        self.expansion_cite = None;
2029        self.expanding_library_body = false;
2030    }
2031
2032    fn record_numeric_effects(
2033        &mut self,
2034        id: NodeId,
2035        op: &str,
2036        effect: &OpEffect,
2037        operands: &[NodeId],
2038        span: SrcSpan,
2039    ) {
2040        let library_cite = self.expansion_cite.map(str::to_string);
2041        if let Some(operand) = required_operand(op, effect.requires, operands) {
2042            let producer_domain = self.graph.node(operand).domain;
2043            if let Some(confidence) = domain_violation(effect.requires, producer_domain) {
2044                let proven = matches!(confidence, DomainViolationConfidence::Proven);
2045                let mut message = format!(
2046                    "`{}` requires {:?} operand, but producer domain is {:?}{}",
2047                    effect.name,
2048                    effect.requires,
2049                    producer_domain,
2050                    if proven {
2051                        " with no discharging guard"
2052                    } else {
2053                        " (producer domain unknown)"
2054                    }
2055                );
2056                if let Some(cite) = &library_cite {
2057                    message.push_str(&format!(" [expanded from {cite}]"));
2058                }
2059                self.graph
2060                    .numeric_domain_violations
2061                    .push(NumericDomainViolation {
2062                        edge_or_node: id,
2063                        op: effect.name.clone(),
2064                        requires: format!("{:?}", effect.requires),
2065                        producer_domain: format!("{producer_domain:?}"),
2066                        proven,
2067                        impact: NumericImpact::LocalOnly,
2068                        span,
2069                        message,
2070                        library_cite: library_cite.clone(),
2071                    });
2072            }
2073        }
2074
2075        let bare = op.rsplit("::").next().unwrap_or(op);
2076        if matches!(bare, "mul" | "broadcast_mul") && operands.len() >= 2 {
2077            if let Some(mut message) =
2078                zero_times_infinity_message(&self.graph, operands[0], operands[1])
2079            {
2080                if let Some(cite) = &library_cite {
2081                    message.push_str(&format!(" [expanded from {cite}]"));
2082                }
2083                self.graph.zero_times_infinity.push(ZeroTimesInfinity {
2084                    edge_or_node: id,
2085                    impact: NumericImpact::LocalOnly,
2086                    span,
2087                    message,
2088                    library_cite,
2089                });
2090            }
2091        }
2092    }
2093
2094    fn result_domain(&self, effect: &OpEffect, op: &str, operands: &[NodeId]) -> NumericDomain {
2095        let bare = op.rsplit("::").next().unwrap_or(op);
2096        if bare == "affine" {
2097            let operand = operands
2098                .first()
2099                .map(|id| self.graph.node(*id).domain)
2100                .unwrap_or(NumericDomain::Unknown);
2101            // Method form: receiver, mul, add — epsilon guard needs mul > 0 and add > 0.
2102            let mul = operands
2103                .get(1)
2104                .and_then(|id| literal_f64(&self.graph.node(*id).kind));
2105            let add = operands
2106                .get(2)
2107                .and_then(|id| literal_f64(&self.graph.node(*id).kind));
2108            return affine_domain(operand, mul, add);
2109        }
2110        match bare {
2111            "mul" | "broadcast_mul" | "add" | "broadcast_add" if operands.len() >= 2 => {
2112                op_semantics::join_domain(
2113                    self.graph.node(operands[0]).domain,
2114                    self.graph.node(operands[1]).domain,
2115                )
2116            }
2117            _ => effect.domain,
2118        }
2119    }
2120}
2121
2122fn required_operand(op: &str, requires: DomainRequirement, operands: &[NodeId]) -> Option<NodeId> {
2123    if matches!(requires, DomainRequirement::None) || operands.is_empty() {
2124        return None;
2125    }
2126    let bare = op.rsplit("::").next().unwrap_or(op);
2127    // For division the non-zero requirement applies to the divisor.
2128    if matches!(bare, "div" | "broadcast_div") && operands.len() >= 2 {
2129        return Some(operands[1]);
2130    }
2131    Some(operands[0])
2132}
2133
2134fn literal_f64(kind: &NodeKind) -> Option<f64> {
2135    let NodeKind::Literal { text } = kind else {
2136        return None;
2137    };
2138    let trimmed = text.trim().trim_end_matches('_');
2139    let trimmed = trimmed
2140        .strip_suffix("f64")
2141        .or_else(|| trimmed.strip_suffix("f32"))
2142        .or_else(|| trimmed.strip_suffix("f16"))
2143        .unwrap_or(trimmed);
2144    trimmed.parse::<f64>().ok()
2145}
2146
2147fn is_undischarged_log(graph: &ExprGraph, node: NodeId) -> bool {
2148    let node = graph.node(node);
2149    let NodeKind::Call { callee } = &node.kind else {
2150        return false;
2151    };
2152    let bare = callee.rsplit("::").next().unwrap_or(callee.as_str());
2153    if bare != "log" {
2154        return false;
2155    }
2156    let Some(operand) = graph
2157        .edges
2158        .iter()
2159        .find(|edge| edge.to == node.id && edge.label.as_deref() != Some("module"))
2160        .map(|edge| edge.from)
2161    else {
2162        return false;
2163    };
2164    domain_violation(
2165        DomainRequirement::StrictlyPositive,
2166        graph.node(operand).domain,
2167    )
2168    .is_some()
2169}
2170
2171fn zero_times_infinity_message(graph: &ExprGraph, left: NodeId, right: NodeId) -> Option<String> {
2172    let left_zero = domain_includes_zero(graph.node(left).domain);
2173    let right_zero = domain_includes_zero(graph.node(right).domain);
2174    let left_log = is_undischarged_log(graph, left);
2175    let right_log = is_undischarged_log(graph, right);
2176    if (left_zero && right_log) || (right_zero && left_log) {
2177        Some(
2178            "multiply combines a domain that includes 0 with an undischarged `log`, \
2179             which yields `0 * -inf = NaN` rather than a loud `-inf` loss"
2180                .to_string(),
2181        )
2182    } else {
2183        None
2184    }
2185}
2186
2187fn transfer(
2188    effect: &OpEffect,
2189    operands: &[NodeId],
2190    explicit: Option<AbstractDtype>,
2191    lookup: impl Fn(NodeId) -> (AbstractDtype, GradState),
2192) -> (AbstractDtype, GradState, EdgeKind) {
2193    let dtypes: Vec<AbstractDtype> = operands.iter().map(|id| lookup(*id).0).collect();
2194    let grads: Vec<GradState> = operands.iter().map(|id| lookup(*id).1).collect();
2195
2196    let dtype = match effect.dtype {
2197        DtypeRule::Preserve => dtypes.first().copied().unwrap_or(AbstractDtype::Unknown),
2198        DtypeRule::SameAsInputs => {
2199            let known: Vec<_> = dtypes.iter().copied().filter(|d| d.is_known()).collect();
2200            if known.is_empty() {
2201                AbstractDtype::Unknown
2202            } else if known.iter().all(|d| *d == known[0]) {
2203                known[0]
2204            } else {
2205                AbstractDtype::Unknown
2206            }
2207        }
2208        DtypeRule::Explicit => explicit.unwrap_or(AbstractDtype::Unknown),
2209        DtypeRule::Fixed(dtype) => dtype,
2210        DtypeRule::Unknown => AbstractDtype::Unknown,
2211    };
2212
2213    let (grad, edge_kind) = match effect.grad {
2214        GradFlow::Severs => (GradState::Severed, EdgeKind::Severing),
2215        GradFlow::LayoutDependent => (GradState::LayoutDependent, EdgeKind::Data),
2216        GradFlow::Propagates => (propagate_grad(&grads), EdgeKind::Data),
2217        GradFlow::Unknown => (GradState::Unknown, EdgeKind::Data),
2218    };
2219
2220    // Once severed, stay severed even if rule said propagates — handled by edge kind on inputs.
2221    let grad = if grads.iter().any(|g| matches!(g, GradState::Severed))
2222        && matches!(effect.grad, GradFlow::Propagates)
2223    {
2224        // Inputs already severed: result is severed for connectivity purposes.
2225        GradState::Severed
2226    } else {
2227        grad
2228    };
2229
2230    (dtype, grad, edge_kind)
2231}
2232
2233fn propagate_grad(grads: &[GradState]) -> GradState {
2234    if grads.is_empty() {
2235        return GradState::Unknown;
2236    }
2237    if grads.iter().any(|g| matches!(g, GradState::Severed)) {
2238        return GradState::Severed;
2239    }
2240    if grads
2241        .iter()
2242        .any(|g| matches!(g, GradState::LayoutDependent))
2243    {
2244        return GradState::LayoutDependent;
2245    }
2246    if grads
2247        .iter()
2248        .any(|g| matches!(g, GradState::Trainable | GradState::Differentiable))
2249    {
2250        return GradState::Differentiable;
2251    }
2252    if grads.iter().all(|g| matches!(g, GradState::Frozen)) {
2253        return GradState::Frozen;
2254    }
2255    GradState::Unknown
2256}
2257
2258fn first_dtype_mismatch(known: &[AbstractDtype]) -> Option<(AbstractDtype, AbstractDtype)> {
2259    let first = *known.first()?;
2260    known
2261        .iter()
2262        .copied()
2263        .find(|d| *d != first)
2264        .map(|other| (first, other))
2265}
2266
2267fn operand_label(op: &str, index: usize, len: usize) -> String {
2268    if len == 2 {
2269        return if index == 0 {
2270            "lhs".into()
2271        } else {
2272            "rhs".into()
2273        };
2274    }
2275    if index == 0
2276        && !matches!(
2277            op,
2278            "cross_entropy" | "nll" | "mse" | "huber" | "binary_cross_entropy_with_logit"
2279        )
2280    {
2281        // method receiver
2282        return "self".into();
2283    }
2284    format!("arg{index}")
2285}
2286
2287fn join_dtypes(ds: &[AbstractDtype]) -> AbstractDtype {
2288    let known: Vec<_> = ds.iter().copied().filter(|d| d.is_known()).collect();
2289    if known.is_empty() {
2290        AbstractDtype::Unknown
2291    } else if known.iter().all(|d| *d == known[0]) {
2292        known[0]
2293    } else {
2294        AbstractDtype::Unknown
2295    }
2296}
2297
2298fn join_grads(gs: &[GradState]) -> GradState {
2299    if gs.is_empty() {
2300        return GradState::Unknown;
2301    }
2302    let first = gs[0];
2303    if gs.iter().all(|g| *g == first) {
2304        first
2305    } else {
2306        GradState::Unknown
2307    }
2308}
2309
2310fn hint_param_grad(type_text: &str, incoming: GradState) -> GradState {
2311    if !matches!(incoming, GradState::Unknown) {
2312        return incoming;
2313    }
2314    match source_type_base(type_text).as_deref() {
2315        Some("Var" | "candle_core::Var") => GradState::Trainable,
2316        _ => GradState::Unknown,
2317    }
2318}
2319
2320fn source_type_base(text: &str) -> Option<String> {
2321    let ty = syn::parse_str::<syn::Type>(text).ok()?;
2322    innermost_type_name(&ty)
2323}
2324
2325fn resolved_return_type(function: &ImplFn, owner_type: &str) -> Option<String> {
2326    let ty = syn::parse_str::<syn::Type>(&function.return_type).ok()?;
2327    let base = result_inner_base(&ty)?;
2328    if base == "Self" {
2329        Some(owner_type.to_string())
2330    } else {
2331        Some(base)
2332    }
2333}
2334
2335fn result_inner_base(ty: &syn::Type) -> Option<String> {
2336    match ty {
2337        syn::Type::Reference(reference) => result_inner_base(&reference.elem),
2338        syn::Type::Paren(paren) => result_inner_base(&paren.elem),
2339        syn::Type::Group(group) => result_inner_base(&group.elem),
2340        syn::Type::Path(path) => {
2341            let segment = path.path.segments.last()?;
2342            if matches!(
2343                segment.ident.to_string().as_str(),
2344                "Result" | "Option" | "Box" | "Arc"
2345            ) {
2346                let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments else {
2347                    return None;
2348                };
2349                return arguments.args.iter().find_map(|argument| match argument {
2350                    syn::GenericArgument::Type(inner) => result_inner_base(inner),
2351                    _ => None,
2352                });
2353            }
2354            Some(
2355                path.path
2356                    .segments
2357                    .iter()
2358                    .map(|segment| segment.ident.to_string())
2359                    .collect::<Vec<_>>()
2360                    .join("::"),
2361            )
2362        }
2363        _ => None,
2364    }
2365}
2366
2367fn is_tensor_constructor_name(name: &str) -> bool {
2368    matches!(
2369        name,
2370        "zeros" | "ones" | "from_vec" | "new" | "arange" | "arange_step" | "rand" | "randn"
2371        | "zeros_like" | "ones_like"
2372    )
2373}
2374
2375fn is_candle_tensor_receiver(type_name: &str) -> bool {
2376    matches!(
2377        type_name.trim_start_matches('&').trim().rsplit("::").next(),
2378        Some("Tensor" | "Var")
2379    )
2380}
2381
2382fn is_non_tensor_tensor_method(method: &str) -> bool {
2383    matches!(
2384        method.rsplit("::").next().unwrap_or(method),
2385        "device"
2386            | "dtype"
2387            | "layout"
2388            | "rank"
2389            | "dims"
2390            | "dim"
2391            | "dims1"
2392            | "dims2"
2393            | "dims3"
2394            | "dims4"
2395            | "dims5"
2396            | "elem_count"
2397            | "stride"
2398            | "is_contiguous"
2399            | "to_scalar"
2400            | "to_vec0"
2401            | "to_vec1"
2402            | "to_vec2"
2403            | "to_vec3"
2404    )
2405}
2406
2407fn normalize_qualified_segments(segments: &[String]) -> String {
2408    segments
2409        .iter()
2410        .map(String::as_str)
2411        .skip_while(|segment| matches!(*segment, "crate" | "self"))
2412        .collect::<Vec<_>>()
2413        .join("::")
2414}
2415
2416fn is_explicit_candle_path(segments: &[String]) -> bool {
2417    matches!(
2418        segments.first().map(String::as_str),
2419        Some(
2420            "candle" | "candle_core" | "candle_nn" | "candle_transformers" | "nn" | "ops" | "loss"
2421        )
2422    )
2423}
2424
2425fn is_scalar_literal(expr: &syn::Expr) -> bool {
2426    match expr {
2427        syn::Expr::Lit(literal) => matches!(
2428            literal.lit,
2429            syn::Lit::Int(_)
2430                | syn::Lit::Float(_)
2431                | syn::Lit::Bool(_)
2432                | syn::Lit::Byte(_)
2433                | syn::Lit::Char(_)
2434        ),
2435        syn::Expr::Unary(unary) => is_scalar_literal(&unary.expr),
2436        syn::Expr::Paren(paren) => is_scalar_literal(&paren.expr),
2437        syn::Expr::Group(group) => is_scalar_literal(&group.expr),
2438        _ => false,
2439    }
2440}
2441
2442fn innermost_type_name(ty: &syn::Type) -> Option<String> {
2443    match ty {
2444        syn::Type::Reference(reference) => innermost_type_name(&reference.elem),
2445        syn::Type::Paren(paren) => innermost_type_name(&paren.elem),
2446        syn::Type::Group(group) => innermost_type_name(&group.elem),
2447        syn::Type::Path(path) => {
2448            let segment = path.path.segments.last()?;
2449            if let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments {
2450                if let Some(inner) = arguments.args.iter().find_map(|argument| match argument {
2451                    syn::GenericArgument::Type(inner) => Some(inner),
2452                    _ => None,
2453                }) {
2454                    return innermost_type_name(inner);
2455                }
2456            }
2457            Some(
2458                path.path
2459                    .segments
2460                    .iter()
2461                    .map(|segment| segment.ident.to_string())
2462                    .collect::<Vec<_>>()
2463                    .join("::"),
2464            )
2465        }
2466        syn::Type::Tuple(tuple) if tuple.elems.len() == 1 => {
2467            tuple.elems.first().and_then(innermost_type_name)
2468        }
2469        _ => None,
2470    }
2471}
2472
2473fn resolve_entrypoint<'a>(
2474    krate: &'a Crate,
2475    entrypoint: &str,
2476) -> anyhow::Result<(&'a ImplFn, String)> {
2477    let function_candidates = krate.function_candidates(entrypoint);
2478    match function_candidates.as_slice() {
2479        [func] => return Ok((func, String::new())),
2480        [] => {}
2481        _ => {
2482            return Err(anyhow::anyhow!(
2483                "entrypoint `{entrypoint}` is ambiguous ({} free functions)",
2484                function_candidates.len()
2485            ))
2486        }
2487    }
2488    if let Some((ty, method)) = entrypoint.rsplit_once("::") {
2489        let candidates = krate.method_candidates(ty, method);
2490        return match candidates.as_slice() {
2491            [func] => Ok((func, ty.to_string())),
2492            [] => Err(anyhow::anyhow!("entrypoint `{entrypoint}` not found")),
2493            _ => Err(anyhow::anyhow!(
2494                "entrypoint `{entrypoint}` is ambiguous ({} methods); select active Cargo cfg",
2495                candidates.len()
2496            )),
2497        };
2498    }
2499    // Bare method name unique across all loaded impls?
2500    let hits: Vec<_> = krate
2501        .all_methods()
2502        .filter(|func| func.fn_name == entrypoint)
2503        .collect();
2504    match hits.len() {
2505        1 => {
2506            let func = hits[0];
2507            Ok((func, func.qualified_type_name.clone()))
2508        }
2509        0 => Err(anyhow::anyhow!("entrypoint `{entrypoint}` not found")),
2510        _ => Err(anyhow::anyhow!(
2511            "entrypoint `{entrypoint}` is ambiguous; use Type::method"
2512        )),
2513    }
2514}
2515
2516fn bind_pat(env: &mut HashMap<String, NodeId>, pat: &syn::Pat, value: NodeId) {
2517    match pat {
2518        syn::Pat::Ident(id) => {
2519            env.insert(id.ident.to_string(), value);
2520        }
2521        syn::Pat::Type(t) => bind_pat(env, &t.pat, value),
2522        syn::Pat::Reference(r) => bind_pat(env, &r.pat, value),
2523        syn::Pat::Tuple(t) => {
2524            // Without splitting the value node, bind each subpat to the same node (honest Unknown
2525            // would be worse for simple `(a, b) = …` where we lack tuple projection).
2526            for p in &t.elems {
2527                bind_pat(env, p, value);
2528            }
2529        }
2530        syn::Pat::TupleStruct(t) => {
2531            for p in &t.elems {
2532                bind_pat(env, p, value);
2533            }
2534        }
2535        syn::Pat::Slice(s) => {
2536            for p in &s.elems {
2537                bind_pat(env, p, value);
2538            }
2539        }
2540        syn::Pat::Wild(_) => {}
2541        _ => {}
2542    }
2543}
2544
2545fn path_ident(p: &syn::ExprPath) -> Option<String> {
2546    if p.path.segments.len() == 1 {
2547        Some(p.path.segments[0].ident.to_string())
2548    } else {
2549        None
2550    }
2551}
2552
2553fn path_last(path: &syn::Path) -> String {
2554    path.segments
2555        .last()
2556        .map(|s| s.ident.to_string())
2557        .unwrap_or_default()
2558}
2559
2560fn path_text(path: &syn::Path) -> String {
2561    path.segments
2562        .iter()
2563        .map(|s| s.ident.to_string())
2564        .collect::<Vec<_>>()
2565        .join("::")
2566}
2567
2568fn lit_text(lit: &syn::Lit) -> String {
2569    match lit {
2570        syn::Lit::Str(s) => s.value(),
2571        syn::Lit::Int(i) => i.to_string(),
2572        syn::Lit::Float(f) => f.to_string(),
2573        syn::Lit::Bool(b) => b.value.to_string(),
2574        _ => load::type_text(lit),
2575    }
2576}
2577
2578fn span_of(file: usize, span: proc_macro2::Span) -> SrcSpan {
2579    load::span_of(file, span)
2580}
2581
2582fn span_binop(op: &syn::BinOp) -> proc_macro2::Span {
2583    op.span()
2584}
2585
2586fn expr_kind_name(expr: &syn::Expr) -> &'static str {
2587    match expr {
2588        syn::Expr::Array(_) => "array",
2589        syn::Expr::Async(_) => "async",
2590        syn::Expr::Cast(_) => "cast",
2591        syn::Expr::Index(_) => "index",
2592        syn::Expr::Range(_) => "range",
2593        syn::Expr::Struct(_) => "struct",
2594        syn::Expr::Repeat(_) => "repeat",
2595        syn::Expr::Unsafe(_) => "unsafe",
2596        syn::Expr::While(_) => "while",
2597        syn::Expr::Loop(_) => "loop",
2598        syn::Expr::ForLoop(_) => "for",
2599        syn::Expr::Break(_) => "break",
2600        syn::Expr::Continue(_) => "continue",
2601        syn::Expr::Yield(_) => "yield",
2602        _ => "expr",
2603    }
2604}
2605
2606fn dtype_from_syn_expr(expression: &syn::Expr) -> Option<AbstractDtype> {
2607    let expression = strip_dataflow_expr(expression);
2608    if let syn::Expr::Path(path) = expression {
2609        let segments: Vec<_> = path.path.segments.iter().collect();
2610        if segments.len() >= 2 && segments[segments.len() - 2].ident == "DType" {
2611            return Some(AbstractDtype::parse(
2612                &segments[segments.len() - 1].ident.to_string(),
2613            ));
2614        }
2615    }
2616    None
2617}
2618
2619fn dtype_from_scalar_expr(expression: &syn::Expr) -> Option<AbstractDtype> {
2620    match strip_dataflow_expr(expression) {
2621        syn::Expr::Lit(literal) => match &literal.lit {
2622            syn::Lit::Float(value) => Some(match value.suffix() {
2623                "" | "f64" => AbstractDtype::F64,
2624                "f32" => AbstractDtype::F32,
2625                _ => AbstractDtype::Unknown,
2626            }),
2627            syn::Lit::Int(value) => Some(match value.suffix() {
2628                "" | "i64" => AbstractDtype::I64,
2629                "i32" => AbstractDtype::I32,
2630                "u32" => AbstractDtype::U32,
2631                "u8" => AbstractDtype::U8,
2632                _ => AbstractDtype::Unknown,
2633            }),
2634            _ => None,
2635        },
2636        other => dtype_from_syn_expr(other),
2637    }
2638}
2639
2640fn dtype_from_collection_expr(expression: &syn::Expr) -> Option<AbstractDtype> {
2641    match strip_dataflow_expr(expression) {
2642        syn::Expr::Array(array) => array.elems.first().and_then(dtype_from_scalar_expr),
2643        syn::Expr::Repeat(repeat) => dtype_from_scalar_expr(&repeat.expr),
2644        other => dtype_from_scalar_expr(other),
2645    }
2646}
2647
2648fn strip_dataflow_expr(expression: &syn::Expr) -> &syn::Expr {
2649    match expression {
2650        syn::Expr::Group(group) => strip_dataflow_expr(&group.expr),
2651        syn::Expr::Paren(paren) => strip_dataflow_expr(&paren.expr),
2652        syn::Expr::Reference(reference) => strip_dataflow_expr(&reference.expr),
2653        syn::Expr::Try(value) => strip_dataflow_expr(&value.expr),
2654        other => other,
2655    }
2656}