1use 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
19pub const CONTRACT_SCHEMA: &str = "candle-graph/contracts/1";
21
22#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
24pub struct ContractAnalysis {
25 pub schema_version: String,
26 pub functions: Vec<FunctionContracts>,
27}
28
29#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
31pub struct FunctionContracts {
32 pub qualified_name: String,
33 pub tensors: Vec<TensorContract>,
34}
35
36pub 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
55pub fn analyze_function(krate: &Crate, function: &ImplFn) -> FunctionContracts {
57 FunctionAnalyzer::new(krate, function).run()
58}
59
60pub 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
72pub 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 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 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}