Skip to main content

candle_graph/
contracts.rs

1//! Conservative tensor-contract inference from qualified Rust functions.
2//!
3//! This pass intentionally does not try to become a Rust type checker. It starts from
4//! [`crate::load::ImplFn`] signatures, follows local tensor expressions, and records only facts
5//! justified by syntax or a small set of Candle operations. Unknown values stay unknown.
6
7use std::collections::{BTreeMap, HashMap};
8
9use quote::ToTokens;
10use serde::{Deserialize, Serialize};
11use syn::spanned::Spanned;
12
13use crate::load::{Crate, ImplFn};
14use crate::model_ir::{
15    Confidence, DeviceFact, Dimension, Evidence, EvidenceKind, LayoutFact, ShapeFact, StableId,
16    TensorContract, TensorRole,
17};
18
19/// Schema identifier for the standalone contract pass.
20pub const CONTRACT_SCHEMA: &str = "candle-graph/contracts/1";
21
22/// Contract facts grouped by their fully-qualified owner function.
23#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
24pub struct ContractAnalysis {
25    pub schema_version: String,
26    pub functions: Vec<FunctionContracts>,
27}
28
29/// Tensor facts found in one function.
30#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
31pub struct FunctionContracts {
32    pub qualified_name: String,
33    pub tensors: Vec<TensorContract>,
34}
35
36/// Analyze every uniquely indexed free function and inherent method.
37pub fn analyze(krate: &Crate) -> ContractAnalysis {
38    let mut functions: BTreeMap<String, &ImplFn> = BTreeMap::new();
39    for function in krate.qualified_functions.values() {
40        functions.insert(function.qualified_name.clone(), function);
41    }
42    for function in krate.qualified_methods.values() {
43        functions.insert(function.qualified_name.clone(), function);
44    }
45
46    ContractAnalysis {
47        schema_version: CONTRACT_SCHEMA.to_string(),
48        functions: functions
49            .into_values()
50            .map(|function| analyze_function(krate, function))
51            .collect(),
52    }
53}
54
55/// Analyze one already-resolved function.
56pub fn analyze_function(krate: &Crate, function: &ImplFn) -> FunctionContracts {
57    FunctionAnalyzer::new(krate, function).run()
58}
59
60/// Analyze all definitions matching a bare or fully-qualified free-function name.
61///
62/// Bare-name collisions are retained. Callers that need exactly one result should supply the
63/// qualified name and verify the returned length.
64pub fn functions_named(krate: &Crate, name: &str) -> Vec<FunctionContracts> {
65    krate
66        .function_candidates(name)
67        .into_iter()
68        .map(|function| analyze_function(krate, function))
69        .collect()
70}
71
72/// Analyze all inherent methods matching a bare or qualified owner.
73pub fn methods_named(krate: &Crate, owner: &str, method: &str) -> Vec<FunctionContracts> {
74    krate
75        .method_candidates(owner, method)
76        .into_iter()
77        .map(|function| analyze_function(krate, function))
78        .collect()
79}
80
81#[derive(Clone)]
82struct Fact {
83    shape: ShapeFact,
84    dtype: String,
85    device: DeviceFact,
86    layout: LayoutFact,
87    requires_grad: Option<bool>,
88    evidence: Vec<Evidence>,
89}
90
91impl Default for Fact {
92    fn default() -> Self {
93        Self {
94            shape: ShapeFact::default(),
95            dtype: "unknown".to_string(),
96            device: DeviceFact::Unknown,
97            layout: LayoutFact::Unknown,
98            requires_grad: None,
99            evidence: Vec::new(),
100        }
101    }
102}
103
104struct FunctionAnalyzer<'a> {
105    krate: &'a Crate,
106    function: &'a ImplFn,
107    facts: BTreeMap<String, (TensorRole, Fact)>,
108    shapes: HashMap<String, ShapeFact>,
109    output_count: usize,
110}
111
112impl<'a> FunctionAnalyzer<'a> {
113    fn new(krate: &'a Crate, function: &'a ImplFn) -> Self {
114        Self {
115            krate,
116            function,
117            facts: BTreeMap::new(),
118            shapes: HashMap::new(),
119            output_count: 0,
120        }
121    }
122
123    fn run(mut self) -> FunctionContracts {
124        for (name, ty) in self
125            .function
126            .params
127            .iter()
128            .zip(self.function.param_types.iter())
129        {
130            if name != "self" && type_contains_tensor(ty) {
131                let mut fact = Fact::default();
132                fact.evidence.push(self.evidence(
133                    self.function.span.line,
134                    Confidence::Proven,
135                    format!("parameter `{name}` has Tensor type `{ty}`"),
136                ));
137                self.facts.insert(name.clone(), (TensorRole::Input, fact));
138            }
139        }
140
141        self.observe_dimension_bindings(&self.function.block);
142        self.process_block(&self.function.block, true);
143
144        let owner = StableId::new("function", [&self.function.qualified_name]);
145        let tensors = self
146            .facts
147            .into_iter()
148            .map(|(name, (role, fact))| TensorContract {
149                id: StableId::new("tensor", [&self.function.qualified_name, &name]),
150                name,
151                role,
152                owner_function: owner.clone(),
153                parameter: None,
154                shape: fact.shape,
155                dtype: fact.dtype,
156                device: fact.device,
157                layout: fact.layout,
158                requires_grad: fact.requires_grad,
159                evidence: fact.evidence,
160            })
161            .collect();
162
163        FunctionContracts {
164            qualified_name: self.function.qualified_name.clone(),
165            tensors,
166        }
167    }
168
169    fn evidence(&self, line: usize, confidence: Confidence, detail: impl Into<String>) -> Evidence {
170        let source = self
171            .krate
172            .files
173            .get(self.function.span.file)
174            .map(|file| format!("{}:{line}", file.rel));
175        Evidence {
176            kind: EvidenceKind::Source,
177            confidence,
178            source,
179            detail: detail.into(),
180        }
181    }
182
183    fn expr_evidence(
184        &self,
185        expr: &syn::Expr,
186        confidence: Confidence,
187        detail: impl Into<String>,
188    ) -> Evidence {
189        self.evidence(expr.span().start().line, confidence, detail)
190    }
191
192    fn observe_dimension_bindings(&mut self, block: &syn::Block) {
193        for statement in &block.stmts {
194            match statement {
195                syn::Stmt::Local(local) => self.observe_dimension_local(local),
196                syn::Stmt::Expr(expr, _) => self.observe_dimension_expr(expr),
197                syn::Stmt::Item(_) | syn::Stmt::Macro(_) => {}
198            }
199        }
200    }
201
202    fn observe_dimension_expr(&mut self, expr: &syn::Expr) {
203        match strip_expr(expr) {
204            syn::Expr::Block(block) => self.observe_dimension_bindings(&block.block),
205            syn::Expr::If(expr_if) => {
206                self.observe_dimension_bindings(&expr_if.then_branch);
207                if let Some((_, otherwise)) = &expr_if.else_branch {
208                    self.observe_dimension_expr(otherwise);
209                }
210            }
211            syn::Expr::ForLoop(loop_expr) => self.observe_dimension_bindings(&loop_expr.body),
212            syn::Expr::While(loop_expr) => self.observe_dimension_bindings(&loop_expr.body),
213            syn::Expr::Loop(loop_expr) => self.observe_dimension_bindings(&loop_expr.body),
214            syn::Expr::Match(expr_match) => {
215                for arm in &expr_match.arms {
216                    self.observe_dimension_expr(&arm.body);
217                }
218            }
219            _ => {}
220        }
221    }
222
223    fn observe_dimension_local(&mut self, local: &syn::Local) {
224        let Some(init) = &local.init else {
225            return;
226        };
227        let Some((tensor, method, args)) = dimension_call(&init.expr) else {
228            self.observe_dimension_expr(&init.expr);
229            return;
230        };
231        if !self.facts.contains_key(&tensor) {
232            return;
233        }
234
235        let pattern_rank = if method == "dims" {
236            match &local.pat {
237                syn::Pat::Slice(slice) => Some(slice.elems.len()),
238                syn::Pat::Type(typed) => match &*typed.pat {
239                    syn::Pat::Slice(slice) => Some(slice.elems.len()),
240                    _ => None,
241                },
242                _ => None,
243            }
244        } else {
245            dims_rank(&method)
246        };
247        if let Some(rank) = pattern_rank {
248            let names = pattern_names(&local.pat);
249            if names.len() == rank {
250                let shape = ShapeFact {
251                    rank: Some(rank),
252                    dimensions: names.iter().map(|name| dimension(name)).collect(),
253                    source_expr: Some(format!("{tensor}.{method}()")),
254                };
255                let evidence = self.evidence(
256                    local.span().start().line,
257                    Confidence::Proven,
258                    format!(
259                        "`{tensor}.{method}()` destructures into [{}]",
260                        names.join(", ")
261                    ),
262                );
263                if let Some((_, fact)) = self.facts.get_mut(&tensor) {
264                    fact.shape = shape;
265                    fact.evidence.push(evidence);
266                }
267            }
268            return;
269        }
270
271        if method == "dim" {
272            let Some(axis) = args.first().and_then(|expr| literal_usize(expr)) else {
273                return;
274            };
275            let names = pattern_names(&local.pat);
276            let Some(name) = names.first() else {
277                return;
278            };
279            let evidence = self.evidence(
280                local.span().start().line,
281                Confidence::Proven,
282                format!("`{name}` is axis {axis} of `{tensor}`"),
283            );
284            if let Some((_, fact)) = self.facts.get_mut(&tensor) {
285                if fact.shape.rank.is_some_and(|rank| axis < rank) {
286                    while fact.shape.dimensions.len() < fact.shape.rank.unwrap_or(0) {
287                        fact.shape
288                            .dimensions
289                            .push(dimension(&format!("axis_{}", fact.shape.dimensions.len())));
290                    }
291                    fact.shape.dimensions[axis] = dimension(name);
292                }
293                fact.evidence.push(evidence);
294            }
295        }
296    }
297
298    fn process_block(&mut self, block: &syn::Block, capture_tail: bool) {
299        for (index, statement) in block.stmts.iter().enumerate() {
300            let is_tail = capture_tail && index + 1 == block.stmts.len();
301            match statement {
302                syn::Stmt::Local(local) => self.process_local(local),
303                syn::Stmt::Expr(expr, semicolon) => {
304                    self.process_nested(expr);
305                    if is_tail
306                        && semicolon.is_none()
307                        && type_contains_tensor(&self.function.return_type)
308                    {
309                        self.record_outputs(expr);
310                    }
311                }
312                syn::Stmt::Item(_) | syn::Stmt::Macro(_) => {}
313            }
314        }
315    }
316
317    fn process_nested(&mut self, expr: &syn::Expr) {
318        match strip_expr(expr) {
319            syn::Expr::Block(block) => self.process_block(&block.block, false),
320            syn::Expr::If(expr_if) => {
321                self.process_block(&expr_if.then_branch, false);
322                if let Some((_, otherwise)) = &expr_if.else_branch {
323                    self.process_nested(otherwise);
324                }
325            }
326            syn::Expr::ForLoop(loop_expr) => self.process_block(&loop_expr.body, false),
327            syn::Expr::While(loop_expr) => self.process_block(&loop_expr.body, false),
328            syn::Expr::Loop(loop_expr) => self.process_block(&loop_expr.body, false),
329            syn::Expr::Match(expr_match) => {
330                for arm in &expr_match.arms {
331                    self.process_nested(&arm.body);
332                }
333            }
334            syn::Expr::Return(return_expr) => {
335                if let Some(value) = &return_expr.expr {
336                    self.record_outputs(value);
337                }
338            }
339            _ => {}
340        }
341    }
342
343    fn process_local(&mut self, local: &syn::Local) {
344        // The initial dimension pre-pass sees tensor parameters. Re-check while walking so
345        // dimensions of tensor-valued locals declared earlier in the block are also captured.
346        self.observe_dimension_local(local);
347        let Some(init) = &local.init else {
348            return;
349        };
350        self.process_nested(&init.expr);
351
352        let names = pattern_names(&local.pat);
353        if names.len() != 1 {
354            return;
355        }
356        let name = names[0].clone();
357
358        if let Some(shape) = self.shape_expr(&init.expr) {
359            self.shapes.insert(name.clone(), shape);
360        }
361
362        if let Some(mut fact) = self.infer_expr(&init.expr) {
363            fact.evidence.push(self.evidence(
364                local.span().start().line,
365                Confidence::Proven,
366                format!(
367                    "local tensor `{name}` is derived from `{}`",
368                    expr_text(&init.expr)
369                ),
370            ));
371            self.facts.insert(name, (TensorRole::Activation, fact));
372        }
373    }
374
375    fn record_outputs(&mut self, expr: &syn::Expr) {
376        let expr = unwrap_result_expr(expr);
377        if let syn::Expr::Tuple(tuple) = strip_expr(expr) {
378            for element in &tuple.elems {
379                self.record_output(element);
380            }
381        } else {
382            self.record_output(expr);
383        }
384    }
385
386    fn record_output(&mut self, expr: &syn::Expr) {
387        let Some(mut fact) = self.infer_expr(expr) else {
388            return;
389        };
390        let name = if self.output_count == 0 {
391            "return".to_string()
392        } else {
393            format!("return.{}", self.output_count)
394        };
395        self.output_count += 1;
396        fact.evidence.push(self.expr_evidence(
397            expr,
398            Confidence::Proven,
399            format!("function returns tensor expression `{}`", expr_text(expr)),
400        ));
401        self.facts.insert(name, (TensorRole::Output, fact));
402    }
403
404    fn infer_expr(&self, expr: &syn::Expr) -> Option<Fact> {
405        let expr = strip_expr(expr);
406        match expr {
407            syn::Expr::Path(path) => {
408                let name = path.path.get_ident()?.to_string();
409                self.facts.get(&name).map(|(_, fact)| fact.clone())
410            }
411            syn::Expr::MethodCall(call) => self.infer_method(call),
412            syn::Expr::Call(call) => self.infer_call(call),
413            syn::Expr::Binary(binary) => {
414                let mut fact = self
415                    .infer_expr(&binary.left)
416                    .or_else(|| self.infer_expr(&binary.right))?;
417                fact.evidence.push(self.expr_evidence(
418                    expr,
419                    Confidence::Conditional,
420                    format!(
421                        "binary `{}` preserves the selected tensor operand contract",
422                        binary.op.to_token_stream()
423                    ),
424                ));
425                Some(fact)
426            }
427            syn::Expr::Unary(unary) => self.infer_expr(&unary.expr),
428            syn::Expr::Index(index) => self.infer_expr(&index.expr),
429            syn::Expr::If(expr_if) => {
430                let then_fact = expr_if
431                    .then_branch
432                    .stmts
433                    .last()
434                    .and_then(|stmt| match stmt {
435                        syn::Stmt::Expr(expr, None) => self.infer_expr(expr),
436                        _ => None,
437                    });
438                let else_fact = expr_if
439                    .else_branch
440                    .as_ref()
441                    .and_then(|(_, expr)| self.infer_expr(expr));
442                merge_facts(then_fact, else_fact)
443            }
444            syn::Expr::Match(expr_match) => expr_match
445                .arms
446                .iter()
447                .filter_map(|arm| self.infer_expr(&arm.body))
448                .reduce(merge_two_facts),
449            _ => None,
450        }
451    }
452
453    fn infer_method(&self, call: &syn::ExprMethodCall) -> Option<Fact> {
454        let method = call.method.to_string();
455
456        // Extraction methods return host values, not tensors.
457        if method.starts_with("to_vec")
458            || matches!(
459                method.as_str(),
460                "dims"
461                    | "dims1"
462                    | "dims2"
463                    | "dims3"
464                    | "dims4"
465                    | "dims5"
466                    | "dim"
467                    | "dtype"
468                    | "device"
469                    | "rank"
470                    | "elem_count"
471                    | "to_scalar"
472            )
473        {
474            return None;
475        }
476
477        let receiver_name = root_ident(&call.receiver);
478        let mut fact = self.infer_expr(&call.receiver).or_else(|| {
479            if matches!(
480                method.as_str(),
481                "forward" | "forward_t" | "forward_diff" | "apply" | "apply_t"
482            ) {
483                call.args.iter().find_map(|arg| self.infer_expr(arg))
484            } else {
485                None
486            }
487        })?;
488
489        match method.as_str() {
490            "to_dtype" => {
491                if let Some(dtype) = call.args.first() {
492                    fact.dtype = dtype_fact(dtype);
493                    fact.evidence.push(self.expr_evidence(
494                        &syn::Expr::MethodCall(call.clone()),
495                        Confidence::Proven,
496                        format!("`{method}` sets dtype to `{}`", fact.dtype),
497                    ));
498                }
499            }
500            "to_device" => {
501                if let Some(device) = call.args.first() {
502                    fact.device = device_fact(device);
503                    fact.evidence.push(self.expr_evidence(
504                        &syn::Expr::MethodCall(call.clone()),
505                        Confidence::Proven,
506                        format!("`to_device` sets device from `{}`", expr_text(device)),
507                    ));
508                }
509            }
510            "contiguous" => {
511                fact.layout = LayoutFact::Contiguous;
512                fact.evidence.push(self.expr_evidence(
513                    &syn::Expr::MethodCall(call.clone()),
514                    Confidence::Proven,
515                    "`contiguous` materializes contiguous layout",
516                ));
517            }
518            "transpose" | "t" => {
519                if method == "transpose" && call.args.len() == 2 {
520                    if let (Some(left), Some(right)) = (
521                        call.args.first().and_then(literal_usize),
522                        call.args.iter().nth(1).and_then(literal_usize),
523                    ) {
524                        if fact.shape.rank.is_some()
525                            && left < fact.shape.dimensions.len()
526                            && right < fact.shape.dimensions.len()
527                        {
528                            fact.shape.dimensions.swap(left, right);
529                        }
530                    }
531                } else if method == "t" && fact.shape.dimensions.len() >= 2 {
532                    let last = fact.shape.dimensions.len() - 1;
533                    fact.shape.dimensions.swap(last - 1, last);
534                }
535                fact.layout = LayoutFact::Strided;
536                fact.evidence.push(self.expr_evidence(
537                    &syn::Expr::MethodCall(call.clone()),
538                    Confidence::Proven,
539                    format!("`{method}` produces a strided view"),
540                ));
541            }
542            "permute" => {
543                if let Some(order) = call.args.first().and_then(index_list) {
544                    if order.len() == fact.shape.dimensions.len()
545                        && order.iter().all(|axis| *axis < order.len())
546                    {
547                        let prior = fact.shape.dimensions.clone();
548                        fact.shape.dimensions =
549                            order.into_iter().map(|axis| prior[axis].clone()).collect();
550                    }
551                }
552                fact.layout = LayoutFact::Strided;
553                fact.evidence.push(self.expr_evidence(
554                    &syn::Expr::MethodCall(call.clone()),
555                    Confidence::Proven,
556                    "`permute` produces a strided view",
557                ));
558            }
559            "narrow" => {
560                if let (Some(axis), Some(length)) = (
561                    call.args.first().and_then(literal_usize),
562                    call.args.iter().nth(2),
563                ) {
564                    if axis < fact.shape.dimensions.len() {
565                        fact.shape.dimensions[axis] = dimension(&expr_text(length));
566                    }
567                }
568                fact.layout = LayoutFact::Strided;
569                fact.evidence.push(self.expr_evidence(
570                    &syn::Expr::MethodCall(call.clone()),
571                    Confidence::Proven,
572                    "`narrow` preserves rank and creates a strided view",
573                ));
574            }
575            "reshape" | "broadcast_as" => {
576                if let Some(shape) = call.args.first().and_then(|expr| self.shape_expr(expr)) {
577                    fact.shape = shape;
578                }
579                if method == "broadcast_as" {
580                    fact.layout = LayoutFact::Strided;
581                }
582                fact.evidence.push(self.expr_evidence(
583                    &syn::Expr::MethodCall(call.clone()),
584                    Confidence::Proven,
585                    format!("`{method}` applies an explicit symbolic shape"),
586                ));
587            }
588            "detach" | "as_detached_tensor" => {
589                fact.requires_grad = Some(false);
590                fact.evidence.push(self.expr_evidence(
591                    &syn::Expr::MethodCall(call.clone()),
592                    Confidence::Proven,
593                    "explicit detach disables gradient tracking",
594                ));
595            }
596            "clone" | "map_err" | "unwrap" | "expect" | "as_ref" => {}
597            "forward" | "forward_t" | "forward_diff" | "apply" | "apply_t" => {
598                if let Some(name) = receiver_name {
599                    fact.evidence.push(self.expr_evidence(
600                        &syn::Expr::MethodCall(call.clone()),
601                        Confidence::Heuristic,
602                        format!(
603                            "`{name}.{method}` propagates dtype/device from its first tensor input; shape remains conservative"
604                        ),
605                    ));
606                    fact.shape = ShapeFact::default();
607                    fact.layout = LayoutFact::Unknown;
608                }
609            }
610            _ => {
611                fact.evidence.push(self.expr_evidence(
612                    &syn::Expr::MethodCall(call.clone()),
613                    Confidence::Conditional,
614                    format!("`{method}` conservatively preserves receiver metadata"),
615                ));
616            }
617        }
618        Some(fact)
619    }
620
621    fn infer_call(&self, call: &syn::ExprCall) -> Option<Fact> {
622        let name = call_path_name(&call.func)?;
623        let short = name.rsplit("::").next().unwrap_or(&name);
624
625        if matches!(short, "Ok" | "Some") {
626            return call.args.first().and_then(|expr| self.infer_expr(expr));
627        }
628
629        if !name.contains("Tensor") {
630            return None;
631        }
632
633        let mut fact = Fact::default();
634        match short {
635            "zeros" | "ones" => {
636                fact.shape = call
637                    .args
638                    .first()
639                    .and_then(|expr| self.shape_expr(expr))
640                    .unwrap_or_default();
641                if let Some(dtype) = call.args.iter().nth(1) {
642                    fact.dtype = dtype_fact(dtype);
643                }
644                if let Some(device) = call.args.iter().nth(2) {
645                    fact.device = device_fact(device);
646                }
647                fact.layout = LayoutFact::Contiguous;
648                fact.requires_grad = Some(false);
649            }
650            "rand" | "randn" => {
651                fact.shape = call
652                    .args
653                    .iter()
654                    .nth(2)
655                    .and_then(|expr| self.shape_expr(expr))
656                    .unwrap_or_default();
657                if let Some(value) = call.args.first() {
658                    fact.dtype = scalar_dtype(value).unwrap_or_else(|| "unknown".to_string());
659                }
660                if let Some(device) = call.args.iter().nth(3) {
661                    fact.device = device_fact(device);
662                }
663                fact.layout = LayoutFact::Contiguous;
664                fact.requires_grad = Some(false);
665            }
666            "from_vec" => {
667                fact.shape = call
668                    .args
669                    .iter()
670                    .nth(1)
671                    .and_then(|expr| self.shape_expr(expr))
672                    .unwrap_or_default();
673                if let Some(values) = call.args.first() {
674                    fact.dtype = collection_dtype(values).unwrap_or_else(|| "unknown".to_string());
675                }
676                if let Some(device) = call.args.iter().nth(2) {
677                    fact.device = device_fact(device);
678                }
679                fact.layout = LayoutFact::Contiguous;
680                fact.requires_grad = Some(false);
681            }
682            "new" => {
683                if let Some(value) = call.args.first() {
684                    fact.dtype = scalar_dtype(value)
685                        .or_else(|| collection_dtype(value))
686                        .unwrap_or_else(|| "unknown".to_string());
687                }
688                if let Some(device) = call.args.iter().nth(1) {
689                    fact.device = device_fact(device);
690                }
691                fact.layout = LayoutFact::Contiguous;
692                fact.requires_grad = Some(false);
693            }
694            "arange" | "arange_step" => {
695                if let Some(value) = call.args.first() {
696                    fact.dtype = scalar_dtype(value).unwrap_or_else(|| "unknown".to_string());
697                }
698                let device_index = if short == "arange" { 2 } else { 3 };
699                if let Some(device) = call.args.iter().nth(device_index) {
700                    fact.device = device_fact(device);
701                }
702                fact.shape = ShapeFact {
703                    rank: Some(1),
704                    dimensions: vec![dimension("range_len")],
705                    source_expr: Some(expr_text(syn::Expr::Call(call.clone()))),
706                };
707                fact.layout = LayoutFact::Contiguous;
708                fact.requires_grad = Some(false);
709            }
710            _ => return None,
711        }
712        fact.evidence.push(self.expr_evidence(
713            &syn::Expr::Call(call.clone()),
714            Confidence::Proven,
715            format!("Tensor::{short} constructor"),
716        ));
717        Some(fact)
718    }
719
720    fn shape_expr(&self, expr: &syn::Expr) -> Option<ShapeFact> {
721        let expr = strip_expr(expr);
722        if let syn::Expr::Path(path) = expr {
723            if let Some(name) = path.path.get_ident() {
724                if let Some(shape) = self.shapes.get(&name.to_string()) {
725                    return Some(shape.clone());
726                }
727            }
728        }
729
730        let dimensions: Vec<Dimension> = match expr {
731            syn::Expr::Tuple(tuple) => tuple
732                .elems
733                .iter()
734                .map(|expr| dimension(&expr_text(expr)))
735                .collect(),
736            syn::Expr::Array(array) => array
737                .elems
738                .iter()
739                .map(|expr| dimension(&expr_text(expr)))
740                .collect(),
741            syn::Expr::Macro(mac) if mac.mac.path.is_ident("vec") => {
742                let parser =
743                    syn::punctuated::Punctuated::<syn::Expr, syn::Token![,]>::parse_terminated;
744                use syn::parse::Parser;
745                parser
746                    .parse2(mac.mac.tokens.clone())
747                    .ok()?
748                    .iter()
749                    .map(|expr| dimension(&expr_text(expr)))
750                    .collect()
751            }
752            _ => return None,
753        };
754        Some(ShapeFact {
755            rank: Some(dimensions.len()),
756            dimensions,
757            source_expr: Some(expr_text(expr)),
758        })
759    }
760}
761
762fn type_contains_tensor(ty: &str) -> bool {
763    ty.split(|ch: char| !(ch.is_ascii_alphanumeric() || ch == '_'))
764        .any(|part| part == "Tensor")
765}
766
767fn strip_expr(mut expr: &syn::Expr) -> &syn::Expr {
768    loop {
769        expr = match expr {
770            syn::Expr::Try(value) => &value.expr,
771            syn::Expr::Await(value) => &value.base,
772            syn::Expr::Paren(value) => &value.expr,
773            syn::Expr::Group(value) => &value.expr,
774            syn::Expr::Reference(value) => &value.expr,
775            _ => return expr,
776        };
777    }
778}
779
780fn unwrap_result_expr(expr: &syn::Expr) -> &syn::Expr {
781    let expr = strip_expr(expr);
782    if let syn::Expr::Call(call) = expr {
783        if matches!(call_path_name(&call.func).as_deref(), Some("Ok" | "Some")) {
784            if let Some(inner) = call.args.first() {
785                return strip_expr(inner);
786            }
787        }
788    }
789    expr
790}
791
792fn root_ident(expr: &syn::Expr) -> Option<String> {
793    match strip_expr(expr) {
794        syn::Expr::Path(path) => path.path.get_ident().map(ToString::to_string),
795        syn::Expr::MethodCall(call) => root_ident(&call.receiver),
796        syn::Expr::Field(field) => root_ident(&field.base),
797        _ => None,
798    }
799}
800
801fn call_path_name(expr: &syn::Expr) -> Option<String> {
802    let syn::Expr::Path(path) = strip_expr(expr) else {
803        return None;
804    };
805    Some(
806        path.path
807            .segments
808            .iter()
809            .map(|part| part.ident.to_string())
810            .collect::<Vec<_>>()
811            .join("::"),
812    )
813}
814
815fn dimension_call(expr: &syn::Expr) -> Option<(String, String, Vec<&syn::Expr>)> {
816    let syn::Expr::MethodCall(call) = strip_expr(expr) else {
817        return None;
818    };
819    let method = call.method.to_string();
820    if matches!(method.as_str(), "unwrap" | "expect" | "map_err") {
821        return dimension_call(&call.receiver);
822    }
823    if !matches!(
824        method.as_str(),
825        "dims" | "dims1" | "dims2" | "dims3" | "dims4" | "dims5" | "dim"
826    ) {
827        return None;
828    }
829    Some((
830        root_ident(&call.receiver)?,
831        method,
832        call.args.iter().collect(),
833    ))
834}
835
836fn dims_rank(method: &str) -> Option<usize> {
837    match method {
838        "dims1" => Some(1),
839        "dims2" => Some(2),
840        "dims3" => Some(3),
841        "dims4" => Some(4),
842        "dims5" => Some(5),
843        _ => None,
844    }
845}
846
847fn pattern_names(pattern: &syn::Pat) -> Vec<String> {
848    match pattern {
849        syn::Pat::Ident(ident) => vec![ident.ident.to_string()],
850        syn::Pat::Tuple(tuple) => tuple.elems.iter().flat_map(pattern_names).collect(),
851        syn::Pat::TupleStruct(tuple) => tuple.elems.iter().flat_map(pattern_names).collect(),
852        syn::Pat::Slice(slice) => slice.elems.iter().flat_map(pattern_names).collect(),
853        syn::Pat::Paren(paren) => pattern_names(&paren.pat),
854        syn::Pat::Type(typed) => pattern_names(&typed.pat),
855        syn::Pat::Reference(reference) => pattern_names(&reference.pat),
856        syn::Pat::Wild(_) => vec!["_".to_string()],
857        _ => Vec::new(),
858    }
859}
860
861fn dimension(expr: &str) -> Dimension {
862    Dimension {
863        name: semantic_dimension(expr),
864        expr: normalize(expr),
865    }
866}
867
868fn semantic_dimension(expr: &str) -> Option<String> {
869    let normalized = normalize(expr).to_ascii_lowercase();
870    let label = match normalized.as_str() {
871        "b" | "batch" | "batch_size" => "batch",
872        "t" | "tokens" | "seq" | "seq_len" | "positions" | "l" => "tokens",
873        "s" | "slots" | "num_slots" => "slots",
874        "a" | "adapter_slots" | "output_slots" | "num_output_slots" => "adapter_slots",
875        "d" | "dim" | "hidden" | "hidden_size" | "model_dim" | "hq" => "hidden",
876        "v" | "vocab" | "vocab_size" => "vocab",
877        "_" => return None,
878        _ => return Some(normalized),
879    };
880    Some(label.to_string())
881}
882
883fn expr_text(value: impl ToTokens) -> String {
884    normalize(&value.to_token_stream().to_string())
885}
886
887fn normalize(text: &str) -> String {
888    text.replace(" :: ", "::")
889        .replace("( ", "(")
890        .replace(" )", ")")
891        .replace("[ ", "[")
892        .replace(" ]", "]")
893        .replace(" ,", ",")
894        .split_whitespace()
895        .collect::<Vec<_>>()
896        .join(" ")
897}
898
899fn dtype_fact(expr: &syn::Expr) -> String {
900    let expr = strip_expr(expr);
901    if let syn::Expr::MethodCall(call) = expr {
902        if call.method == "dtype" {
903            if let Some(name) = root_ident(&call.receiver) {
904                return format!("same_as({name})");
905            }
906        }
907    }
908    let text = expr_text(expr);
909    text.rsplit("::").next().unwrap_or(&text).to_string()
910}
911
912fn device_fact(expr: &syn::Expr) -> DeviceFact {
913    let expr = strip_expr(expr);
914    if let syn::Expr::MethodCall(call) = expr {
915        if call.method == "device" {
916            if let Some(name) = root_ident(&call.receiver) {
917                return DeviceFact::SameAs(name);
918            }
919        }
920    }
921    let text = expr_text(expr);
922    if text.ends_with("Device::Cpu") || text == "Device::Cpu" {
923        DeviceFact::Cpu
924    } else if text.contains("new_cuda") {
925        DeviceFact::Cuda {
926            ordinal: call_last_literal(expr).map(|value| value as u32),
927        }
928    } else if text.contains("new_metal") {
929        DeviceFact::Metal
930    } else {
931        DeviceFact::SameAs(text.trim_start_matches('&').to_string())
932    }
933}
934
935fn call_last_literal(expr: &syn::Expr) -> Option<usize> {
936    match strip_expr(expr) {
937        syn::Expr::Call(call) => call.args.last().and_then(literal_usize),
938        syn::Expr::MethodCall(call) => call.args.last().and_then(literal_usize),
939        _ => None,
940    }
941}
942
943fn scalar_dtype(expr: &syn::Expr) -> Option<String> {
944    match strip_expr(expr) {
945        syn::Expr::Lit(syn::ExprLit {
946            lit: syn::Lit::Float(value),
947            ..
948        }) => Some(
949            match value.suffix() {
950                "" | "f64" => "F64",
951                "f32" => "F32",
952                suffix => suffix,
953            }
954            .to_ascii_uppercase(),
955        ),
956        syn::Expr::Lit(syn::ExprLit {
957            lit: syn::Lit::Int(value),
958            ..
959        }) => Some(
960            match value.suffix() {
961                "" | "i32" => "I32",
962                "i16" => "I16",
963                "u8" => "U8",
964                "u32" => "U32",
965                "i64" => "I64",
966                suffix => suffix,
967            }
968            .to_ascii_uppercase(),
969        ),
970        syn::Expr::Lit(syn::ExprLit {
971            lit: syn::Lit::Bool(_),
972            ..
973        }) => None,
974        _ => None,
975    }
976}
977
978fn collection_dtype(expr: &syn::Expr) -> Option<String> {
979    match strip_expr(expr) {
980        syn::Expr::Array(array) => array.elems.first().and_then(scalar_dtype),
981        syn::Expr::Macro(mac) if mac.mac.path.is_ident("vec") => {
982            use syn::parse::Parser;
983            let parser = syn::punctuated::Punctuated::<syn::Expr, syn::Token![,]>::parse_terminated;
984            parser
985                .parse2(mac.mac.tokens.clone())
986                .ok()?
987                .first()
988                .and_then(scalar_dtype)
989        }
990        _ => None,
991    }
992}
993
994fn literal_usize(expr: &syn::Expr) -> Option<usize> {
995    let syn::Expr::Lit(syn::ExprLit {
996        lit: syn::Lit::Int(value),
997        ..
998    }) = strip_expr(expr)
999    else {
1000        return None;
1001    };
1002    value.base10_parse().ok()
1003}
1004
1005fn index_list(expr: &syn::Expr) -> Option<Vec<usize>> {
1006    match strip_expr(expr) {
1007        syn::Expr::Array(array) => array.elems.iter().map(literal_usize).collect(),
1008        syn::Expr::Tuple(tuple) => tuple.elems.iter().map(literal_usize).collect(),
1009        _ => None,
1010    }
1011}
1012
1013fn merge_facts(left: Option<Fact>, right: Option<Fact>) -> Option<Fact> {
1014    match (left, right) {
1015        (Some(left), Some(right)) => Some(merge_two_facts(left, right)),
1016        (Some(mut fact), None) | (None, Some(mut fact)) => {
1017            fact.shape = ShapeFact::default();
1018            fact.dtype = "unknown".to_string();
1019            fact.device = DeviceFact::Unknown;
1020            fact.layout = LayoutFact::Unknown;
1021            fact.requires_grad = None;
1022            for evidence in &mut fact.evidence {
1023                evidence.confidence = Confidence::Conditional;
1024            }
1025            Some(fact)
1026        }
1027        (None, None) => None,
1028    }
1029}
1030
1031fn merge_two_facts(mut left: Fact, right: Fact) -> Fact {
1032    if left.shape != right.shape {
1033        left.shape = ShapeFact::default();
1034    }
1035    if left.dtype != right.dtype {
1036        left.dtype = "unknown".to_string();
1037    }
1038    if left.device != right.device {
1039        left.device = DeviceFact::Unknown;
1040    }
1041    if left.layout != right.layout {
1042        left.layout = LayoutFact::Unknown;
1043    }
1044    if left.requires_grad != right.requires_grad {
1045        left.requires_grad = None;
1046    }
1047    left.evidence.extend(right.evidence);
1048    left
1049}