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 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 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 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}