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