use std::collections::HashMap;
use crate::ir::*;
use crate::known::{self, ParamKind, PrefixOp};
use crate::load::{Container, Crate};
const MAX_DEPTH: usize = 32;
const MAX_INSTANCES: usize = 20_000;
#[derive(Clone, Debug)]
struct VbVal {
root: String,
key: Key,
certainty: Certainty,
}
#[derive(Clone, Default)]
struct Env {
vb: HashMap<String, VbVal>,
closures: HashMap<String, syn::ExprClosure>,
}
impl Env {
fn bind_param(&mut self, name: String, value: Option<VbVal>) {
match value {
Some(val) => {
self.vb.insert(name, val);
}
None => {
self.vb.remove(&name);
}
}
}
}
pub struct Extractor<'a> {
krate: &'a Crate,
known_candle_constructors: bool,
out: Structure,
defs: HashMap<(String, String), ModuleDefId>,
sites: HashMap<(ModuleDefId, usize, usize, usize, String), ParamSiteId>,
stack: Vec<(String, String)>,
truncated: bool,
}
#[derive(Clone)]
struct Ctx {
def: ModuleDefId,
instance: ModuleInstanceId,
file: usize,
conditional: Option<String>,
repeat: Option<Repeat>,
depth: usize,
}
impl Ctx {
fn certainty(&self) -> Certainty {
match &self.conditional {
Some(reason) => Certainty::Conditional(reason.clone()),
None => Certainty::Certain,
}
}
fn in_branch(&self, reason: impl Into<String>) -> Self {
let mut next = self.clone();
if next.conditional.is_none() {
next.conditional = Some(reason.into());
}
next
}
}
impl<'a> Extractor<'a> {
pub fn new(krate: &'a Crate) -> Self {
Self {
krate,
known_candle_constructors: true,
out: Structure::default(),
defs: HashMap::new(),
sites: HashMap::new(),
stack: Vec::new(),
truncated: false,
}
}
pub fn for_candle_version(krate: &'a Crate, version: Option<&str>) -> Self {
let mut extractor = Self::new(krate);
extractor.known_candle_constructors =
version.is_some_and(crate::op_semantics::is_audited_candle_version);
extractor
}
pub fn run(mut self, root_type: &str, ctor: Option<&str>) -> anyhow::Result<Structure> {
let ctor_name = match ctor {
Some(name) => name.to_string(),
None => self.find_entry_ctor(root_type)?,
};
let candidates: Vec<_> = self
.krate
.method_candidates(root_type, &ctor_name)
.into_iter()
.filter(|func| func.trait_name.is_none())
.collect();
let func = match candidates.as_slice() {
[func] => *func,
[] => {
return Err(anyhow::anyhow!(
"`{root_type}::{ctor_name}` not found in crate"
))
}
_ => {
return Err(anyhow::anyhow!(
"`{root_type}::{ctor_name}` is ambiguous ({} definitions); use a \
module-qualified root and active Cargo cfg",
candidates.len()
))
}
};
let def = self.def_id(root_type, &ctor_name, func.span);
let primary_root = func
.vb_params
.first()
.and_then(|i| func.params.get(*i))
.cloned()
.unwrap_or_else(|| "vb".to_string());
let root = self.out.add_instance(
def,
None,
None,
Key::default(),
primary_root,
false,
None,
func.span,
Certainty::Certain,
);
self.out.root = Some(root);
let mut env = Env::default();
for index in &func.vb_params {
if let Some(name) = func.params.get(*index) {
env.vb.insert(
name.clone(),
VbVal {
root: name.clone(),
key: Key::default(),
certainty: Certainty::Certain,
},
);
}
}
let ctx = Ctx {
def,
instance: root,
file: func.span.file,
conditional: None,
repeat: None,
depth: 0,
};
self.stack.push((root_type.to_string(), ctor_name));
let block = func.block.clone();
self.walk_block(&block, &mut env, &ctx);
self.stack.pop();
if self.truncated {
self.out.diagnose(
SrcSpan::UNKNOWN,
format!("analysis truncated at {MAX_INSTANCES} instances; output is incomplete"),
None,
);
}
self.out.dedupe_params();
self.out.derive_prefixes();
Ok(self.out)
}
fn find_entry_ctor(&self, root_type: &str) -> anyhow::Result<String> {
for preferred in ["new", "load"] {
let candidates: Vec<_> = self
.krate
.method_candidates(root_type, preferred)
.into_iter()
.filter(|func| func.trait_name.is_none() && !func.vb_params.is_empty())
.collect();
match candidates.len() {
0 => {}
1 => return Ok(preferred.to_string()),
count => {
return Err(anyhow::anyhow!(
"`{root_type}::{preferred}` is ambiguous ({count} active-or-cfg-gated \
definitions); select a module-qualified root and Cargo configuration"
))
}
}
}
let mut candidates: Vec<String> = self
.krate
.all_methods()
.filter(|func| {
let owner_matches = if root_type.contains("::") {
func.qualified_type_name == root_type
} else {
func.type_name == root_type
};
owner_matches && func.trait_name.is_none() && !func.vb_params.is_empty()
})
.map(|func| func.fn_name.clone())
.collect();
candidates.sort();
candidates.dedup();
candidates.first().cloned().ok_or_else(|| {
anyhow::anyhow!(
"no constructor taking a VarBuilder found on `{root_type}`; \
pass --ctor to choose one explicitly"
)
})
}
fn def_id(&mut self, type_name: &str, ctor: &str, span: SrcSpan) -> ModuleDefId {
let key = (type_name.to_string(), ctor.to_string());
if let Some(id) = self.defs.get(&key) {
return *id;
}
let id = self
.out
.add_def(type_name.to_string(), Some(ctor.to_string()), span);
self.defs.insert(key, id);
id
}
fn walk_block(&mut self, block: &syn::Block, env: &mut Env, ctx: &Ctx) {
for stmt in &block.stmts {
self.walk_stmt(stmt, env, ctx);
}
}
fn walk_stmt(&mut self, stmt: &syn::Stmt, env: &mut Env, ctx: &Ctx) {
match stmt {
syn::Stmt::Local(local) => {
let Some(init) = &local.init else { return };
if let syn::Expr::Closure(closure) = unwrap_expr(&init.expr) {
if let Some(name) = binding_name(&local.pat) {
env.closures.insert(name, closure.clone());
return;
}
}
if let Some(val) = self.eval_vb(&init.expr, env, ctx) {
if let Some(name) = binding_name(&local.pat) {
env.vb.insert(name, val);
return;
}
}
self.walk_expr(&init.expr, env, ctx);
}
syn::Stmt::Expr(expr, _) => self.walk_expr(expr, env, ctx),
syn::Stmt::Item(_) => {}
syn::Stmt::Macro(m) => {
if macro_cannot_register_params(&m.mac.path) {
return;
}
self.out.diagnose(
crate::load::span_of(ctx.file, m.mac.path.segments[0].ident.span()),
format!(
"macro `{}!` not expanded; any parameters it registers are invisible",
crate::load::type_text(&m.mac.path)
),
None,
);
}
}
}
fn walk_expr(&mut self, expr: &syn::Expr, env: &mut Env, ctx: &Ctx) {
if self.out.instances.len() > MAX_INSTANCES {
self.truncated = true;
return;
}
if self.try_local_closure(expr, env, ctx) {
return;
}
if self.try_param_site(expr, env, ctx) {
return;
}
if self.try_free_function(expr, env, ctx) {
return;
}
if self.try_submodule(expr, env, ctx, None) {
return;
}
match unwrap_expr(expr) {
syn::Expr::Struct(s) => self.walk_struct_literal(s, env, ctx),
syn::Expr::ForLoop(f) => {
let mut inner = ctx.clone();
inner.repeat = Some(Repeat {
var: binding_name(&f.pat).unwrap_or_else(|| "_".to_string()),
bound: crate::load::type_text(&f.expr),
});
let mut scoped = env.clone();
self.walk_block(&f.body, &mut scoped, &inner);
}
syn::Expr::While(w) => {
self.walk_expr(&w.cond, env, ctx);
let mut inner = ctx.clone();
inner.repeat = Some(Repeat {
var: "_".to_string(),
bound: format!("while {}", crate::load::type_text(&w.cond)),
});
let mut scoped = env.clone();
self.walk_block(&w.body, &mut scoped, &inner);
}
syn::Expr::Loop(l) => {
let mut inner = ctx.clone();
inner.repeat = Some(Repeat {
var: "_".to_string(),
bound: "loop".to_string(),
});
let mut scoped = env.clone();
self.walk_block(&l.body, &mut scoped, &inner);
}
syn::Expr::If(i) => {
let branch = ctx.in_branch(format!("if {}", crate::load::type_text(&i.cond)));
let mut scoped = env.clone();
self.walk_block(&i.then_branch, &mut scoped, &branch);
if let Some((_, alt)) = &i.else_branch {
let mut scoped = env.clone();
self.walk_expr(alt, &mut scoped, &branch);
}
}
syn::Expr::Match(m) => {
let branch = ctx.in_branch(format!("match {}", crate::load::type_text(&m.expr)));
for arm in &m.arms {
let mut scoped = env.clone();
self.walk_expr(&arm.body, &mut scoped, &branch);
}
}
syn::Expr::Closure(c) => {
let mut scoped = env.clone();
for param in &c.inputs {
if let Some(name) = binding_name(param) {
scoped.bind_param(name, None);
}
}
self.walk_expr(&c.body, &mut scoped, ctx);
}
syn::Expr::MethodCall(mc) => {
let mut inner = ctx.clone();
if matches!(mc.method.to_string().as_str(), "then" | "then_some") {
inner = inner.in_branch(format!(
"only when {}",
crate::load::type_text(&mc.receiver)
));
}
if mc.method == "map"
&& inner.repeat.is_none()
&& looks_like_iteration(&mc.receiver)
{
inner.repeat = Some(Repeat {
var: "_".to_string(),
bound: crate::load::type_text(&mc.receiver),
});
}
let receiver_vb = self.eval_vb(&mc.receiver, env, &inner);
self.walk_expr(&mc.receiver, env, &inner);
for arg in &mc.args {
if let syn::Expr::Closure(closure) = unwrap_expr(arg) {
let mut scoped = env.clone();
for (index, param) in closure.inputs.iter().enumerate() {
if let Some(name) = binding_name(param) {
let value = if index == 0 {
receiver_vb.clone()
} else {
None
};
scoped.bind_param(name, value);
}
}
self.walk_expr(&closure.body, &mut scoped, &inner);
} else {
self.walk_expr(arg, env, &inner);
}
}
}
other => self.walk_children(other, env, ctx),
}
}
fn try_local_closure(&mut self, expr: &syn::Expr, env: &mut Env, ctx: &Ctx) -> bool {
let syn::Expr::Call(call) = unwrap_expr(expr) else {
return false;
};
let Some(path) = call_path(&call.func) else {
return false;
};
if path.len() != 1 {
return false;
}
let Some(closure) = env.closures.get(&path[0]).cloned() else {
return false;
};
let mut scoped = env.clone();
scoped.closures.remove(&path[0]);
for (param, arg) in closure.inputs.iter().zip(call.args.iter()) {
if let Some(name) = binding_name(param) {
let value = self.eval_vb(arg, env, ctx);
scoped.bind_param(name, value);
}
}
for param in closure.inputs.iter().skip(call.args.len()) {
if let Some(name) = binding_name(param) {
scoped.bind_param(name, None);
}
}
self.walk_expr(&closure.body, &mut scoped, ctx);
true
}
fn walk_struct_literal(&mut self, s: &syn::ExprStruct, env: &mut Env, ctx: &Ctx) {
let type_name = s.path.segments.last().map(|seg| seg.ident.to_string());
let own_type = self.out.def(ctx.def).name.clone();
let group = match &type_name {
Some(name)
if name != "Self"
&& *name != own_type
&& self.krate.struct_candidates(name).len() == 1 =>
{
let def = self.def_id(name, "<struct literal>", span_of_expr(ctx, &s.path));
let id = self.out.add_instance(
def,
Some(ctx.instance),
None,
Key::default(),
String::new(),
true,
ctx.repeat.clone(),
span_of_expr(ctx, &s.path),
ctx.certainty(),
);
Some((def, id))
}
_ => None,
};
let inner = match group {
Some((def, instance)) => Ctx {
def,
instance,
..ctx.clone()
},
None => ctx.clone(),
};
for field in &s.fields {
let name = match &field.member {
syn::Member::Named(id) => Some(id.to_string()),
syn::Member::Unnamed(i) => Some(i.index.to_string()),
};
self.walk_field_expr(&field.expr, env, &inner, name);
}
}
fn walk_field_expr(
&mut self,
expr: &syn::Expr,
env: &mut Env,
ctx: &Ctx,
field: Option<String>,
) {
let ctx = match &field {
Some(name) if self.field_is_option(ctx, name) => {
&ctx.in_branch(format!("Option field `{name}`"))
}
_ => ctx,
};
if self.try_local_closure(expr, env, ctx) {
return;
}
if self.try_param_site(expr, env, ctx) {
return;
}
if self.try_free_function(expr, env, ctx) {
return;
}
if self.try_submodule(expr, env, ctx, field) {
return;
}
self.walk_expr(expr, env, ctx);
}
fn field_is_option(&self, ctx: &Ctx, field: &str) -> bool {
let candidates = self.krate.struct_candidates(&self.out.def(ctx.def).name);
candidates
.first()
.filter(|_| candidates.len() == 1)
.and_then(|s| s.fields.iter().find(|f| f.name == field))
.map(|f| f.ty.container == Container::Option)
.unwrap_or(false)
}
fn walk_children(&mut self, expr: &syn::Expr, env: &mut Env, ctx: &Ctx) {
use syn::Expr as E;
match expr {
E::Call(c) => {
if c.args
.iter()
.any(|arg| self.eval_vb(arg, env, ctx).is_some())
{
let target = call_path(&c.func)
.map(|path| path.join("::"))
.unwrap_or_else(|| crate::load::type_text(&c.func));
self.out.diagnose(
span_of_expr(ctx, c),
format!(
"unresolved call `{target}` receives a VarBuilder; any parameters it \
registers are unknown"
),
None,
);
}
for a in &c.args {
self.walk_expr(a, env, ctx);
}
}
E::Try(t) => self.walk_expr(&t.expr, env, ctx),
E::Reference(r) => self.walk_expr(&r.expr, env, ctx),
E::Paren(p) => self.walk_expr(&p.expr, env, ctx),
E::Group(g) => self.walk_expr(&g.expr, env, ctx),
E::Block(b) => {
let mut scoped = env.clone();
self.walk_block(&b.block, &mut scoped, ctx);
}
E::Unsafe(u) => {
let mut scoped = env.clone();
self.walk_block(&u.block, &mut scoped, ctx);
}
E::Tuple(t) => {
for e in &t.elems {
self.walk_expr(e, env, ctx);
}
}
E::Array(a) => {
for e in &a.elems {
self.walk_expr(e, env, ctx);
}
}
E::Return(r) => {
if let Some(e) = &r.expr {
self.walk_expr(e, env, ctx);
}
}
E::Let(l) => self.walk_expr(&l.expr, env, ctx),
E::Assign(a) => self.walk_expr(&a.right, env, ctx),
E::Binary(b) => {
self.walk_expr(&b.left, env, ctx);
self.walk_expr(&b.right, env, ctx);
}
_ => {}
}
}
fn try_param_site(&mut self, expr: &syn::Expr, env: &mut Env, ctx: &Ctx) -> bool {
let expr = unwrap_expr(expr);
if let syn::Expr::Call(call) = expr {
let Some(path) = call_path(&call.func) else {
return false;
};
let Some(func) = path.last() else {
return false;
};
if path.len() == 1 && env.closures.contains_key(func) {
return false;
}
let resolved_path = self.krate.resolve_unambiguous_import_path(&path);
let Some(ctor) = self
.known_candle_constructors
.then(|| known_constructor(resolved_path.as_slice(), func))
.flatten()
else {
return false;
};
let span = span_of_expr(ctx, expr);
let Some(vb) = self.resolve_vb_arg(&call.args, ctor.vb_arg, env, ctx) else {
self.out.diagnose(
span,
format!(
"`{}` call whose VarBuilder argument could not be resolved to a prefix",
path.join("::")
),
None,
);
return true;
};
let rnn = rnn_config(func, &call.args);
for leaf in ctor.leaves {
if leaf.kind == ParamKind::Bias && rnn.biases == Some(false) {
continue;
}
let shape = constructor_leaf_shape(func, leaf.name, &call.args);
let unconditional = leaf.unconditional
|| (leaf.kind == ParamKind::Bias && rnn.biases == Some(true));
let leaf_certainty = combine(
&vb.certainty,
&ctx.certainty(),
unconditional,
&format!(
"`{func}` registers `{}` only in some configurations",
leaf.name
),
);
let named_by_config = leaf.config_named && !rnn.default_names;
let leaf_seg = if named_by_config {
self.out.diagnose(
span,
format!(
"`{func}` names `{}` from its config argument; unresolved \
`layer_idx`/`direction` leaves a tensor-name family",
leaf.name
),
None,
);
KeySeg::Template {
text: config_named_family(leaf.name),
}
} else {
KeySeg::Literal(leaf.name.to_string())
};
let site = self.site_for(
ctx,
span,
leaf.name,
Acquisition::Constructor {
func: path.join("::"),
cite: ctor.cite,
},
Key::default().push(leaf_seg.clone()),
leaf.kind,
shape,
leaf_certainty.clone(),
);
let key = vb.key.push(leaf_seg);
self.out
.add_param(site, ctx.instance, key, vb.root.clone(), leaf_certainty);
}
return true;
}
if let syn::Expr::MethodCall(mc) = expr {
let method = mc.method.to_string();
let Some(name_arg) = known::raw_get_name_arg(&method) else {
return false;
};
let Some(vb) = self.eval_vb(&mc.receiver, env, ctx) else {
return false;
};
let span = span_of_expr(ctx, expr);
let Some(name_expr) = mc.args.iter().nth(name_arg) else {
return false;
};
let Some(name) = string_literal(name_expr) else {
self.out.diagnose(
span,
format!("`{method}` with a non-literal tensor name; key is unknown"),
Some(vb.key.clone()),
);
return true;
};
let shape = if name_arg > 0 {
mc.args.first().map(crate::load::type_text)
} else {
None
};
let leaf_certainty = combine(&vb.certainty, &ctx.certainty(), true, "");
let site = self.site_for(
ctx,
span,
&name,
Acquisition::RawGet { method },
Key::default().push_literal(&name),
ParamKind::Raw,
shape,
leaf_certainty.clone(),
);
let key = vb.key.push_literal(&name);
self.out
.add_param(site, ctx.instance, key, vb.root, leaf_certainty);
return true;
}
false
}
#[allow(clippy::too_many_arguments)]
fn site_for(
&mut self,
ctx: &Ctx,
span: SrcSpan,
leaf: &str,
acquisition: Acquisition,
relative_key: Key,
kind: ParamKind,
shape: Option<String>,
certainty: Certainty,
) -> ParamSiteId {
let cache_key = (ctx.def, span.file, span.line, span.col, leaf.to_string());
if let Some(id) = self.sites.get(&cache_key) {
return *id;
}
let id = self.out.add_site(
ctx.def,
acquisition,
relative_key,
kind,
shape,
span,
certainty,
);
self.sites.insert(cache_key, id);
id
}
fn try_free_function(&mut self, expr: &syn::Expr, env: &mut Env, ctx: &Ctx) -> bool {
let syn::Expr::Call(call) = unwrap_expr(expr) else {
return false;
};
let Some(path) = call_path(&call.func) else {
return false;
};
let Some(name) = path.last().cloned() else {
return false;
};
let lookup_name = normalize_qualified_path(&path);
let candidates = self.krate.function_candidates(&lookup_name);
let func = match candidates.as_slice() {
[func] => *func,
[] => return false,
_ => {
if candidates.iter().any(|func| !func.vb_params.is_empty()) {
self.out.diagnose(
span_of_expr(ctx, expr),
format!(
"free function `{lookup_name}` is ambiguous ({} definitions); \
parameters cannot be attributed safely",
candidates.len()
),
None,
);
return true;
}
return false;
}
};
if func.vb_params.is_empty() {
return false;
}
let span = span_of_expr(ctx, expr);
let stack_key = ("<free>".to_string(), name.clone());
if ctx.depth >= MAX_DEPTH {
self.out.diagnose(
span,
format!("recursion depth {MAX_DEPTH} exceeded at free function `{name}`"),
None,
);
return true;
}
if self.stack.contains(&stack_key) {
self.out.diagnose(
span,
format!("cycle through free function `{name}`; body not expanded"),
None,
);
return true;
}
let mut helper_env = Env::default();
let mut resolved = 0usize;
for index in &func.vb_params {
let Some(param) = func.params.get(*index) else {
continue;
};
let Some(arg) = call.args.iter().nth(*index) else {
self.out.diagnose(
span,
format!("free function `{name}` is missing VarBuilder argument `{param}`"),
None,
);
continue;
};
match self.eval_vb(arg, env, ctx) {
Some(value) => {
helper_env.vb.insert(param.clone(), value);
resolved += 1;
}
None => self.out.diagnose(
span,
format!(
"free function `{name}` argument `{param}` is a VarBuilder whose prefix \
could not be resolved"
),
None,
),
}
}
if resolved == 0 {
return true;
}
let helper_ctx = Ctx {
file: func.span.file,
depth: ctx.depth + 1,
..ctx.clone()
};
self.stack.push(stack_key);
let block = func.block.clone();
self.walk_block(&block, &mut helper_env, &helper_ctx);
self.stack.pop();
true
}
fn try_submodule(
&mut self,
expr: &syn::Expr,
env: &mut Env,
ctx: &Ctx,
field: Option<String>,
) -> bool {
let expr = unwrap_expr(expr);
let syn::Expr::Call(call) = expr else {
return false;
};
let Some(path) = call_path(&call.func) else {
return false;
};
if path.len() < 2 {
return false;
}
let ctor_name = path[path.len() - 1].clone();
let raw_type = path[..path.len() - 1].join("::");
let type_name = if raw_type == "Self" {
self.out.def(ctx.def).name.clone()
} else {
normalize_qualified_name(&raw_type)
};
let candidates: Vec<_> = self
.krate
.method_candidates(&type_name, &ctor_name)
.into_iter()
.filter(|func| func.trait_name.is_none())
.collect();
let func = match candidates.as_slice() {
[func] => *func,
[] => {
if call
.args
.iter()
.any(|arg| self.eval_vb(arg, env, ctx).is_some())
{
self.out.diagnose(
span_of_expr(ctx, expr),
format!(
"unresolved constructor `{type_name}::{ctor_name}` receives a \
VarBuilder; its parameters are unknown"
),
None,
);
return true;
}
return false;
}
_ => {
self.out.diagnose(
span_of_expr(ctx, expr),
format!(
"constructor `{type_name}::{ctor_name}` is ambiguous ({} definitions); \
subtree not expanded",
candidates.len()
),
None,
);
return true;
}
};
if func.vb_params.is_empty() {
return false;
}
let span = span_of_expr(ctx, expr);
let mut bindings: Vec<(String, VbVal)> = Vec::new();
let mut primary: Option<VbVal> = None;
for index in &func.vb_params {
let Some(arg) = call.args.iter().nth(*index) else {
continue;
};
let Some(name) = func.params.get(*index) else {
continue;
};
match self.eval_vb(arg, env, ctx) {
Some(val) => {
if primary.is_none() {
primary = Some(val.clone());
}
bindings.push((name.clone(), val));
}
None => {
self.out.diagnose(
span,
format!(
"`{type_name}::{ctor_name}` argument `{name}` is a VarBuilder whose \
prefix could not be resolved; parameters under it are missing"
),
None,
);
}
}
}
let Some(primary) = primary else {
self.out.diagnose(
span,
format!("`{type_name}::{ctor_name}` VarBuilder arguments not resolvable"),
None,
);
return true;
};
if ctx.depth >= MAX_DEPTH {
self.out.diagnose(
span,
format!("recursion depth {MAX_DEPTH} exceeded at `{type_name}::{ctor_name}`"),
Some(primary.key),
);
return true;
}
if self.stack.contains(&(type_name.clone(), ctor_name.clone())) {
self.out.diagnose(
span,
format!("cycle through `{type_name}::{ctor_name}`; subtree not expanded"),
Some(primary.key),
);
return true;
}
let certainty = merge(ctx.certainty(), primary.certainty.clone());
if raw_type == "Self" && field.is_none() {
let mut helper_env = Env::default();
for (name, val) in bindings {
helper_env.vb.insert(name, val);
}
let helper_ctx = Ctx {
file: func.span.file,
conditional: match &certainty {
Certainty::Conditional(reason) => Some(reason.clone()),
_ => ctx.conditional.clone(),
},
depth: ctx.depth + 1,
..ctx.clone()
};
self.stack.push((type_name, ctor_name));
let block = func.block.clone();
self.walk_block(&block, &mut helper_env, &helper_ctx);
self.stack.pop();
return true;
}
let child_def = self.def_id(&type_name, &ctor_name, func.span);
let child = self.out.add_instance(
child_def,
Some(ctx.instance),
field,
primary.key.clone(),
primary.root.clone(),
false,
ctx.repeat.clone(),
span,
certainty.clone(),
);
let mut child_env = Env::default();
for (name, val) in bindings {
child_env.vb.insert(name, val);
}
let child_ctx = Ctx {
def: child_def,
instance: child,
file: func.span.file,
conditional: match &certainty {
Certainty::Conditional(reason) => Some(reason.clone()),
_ => None,
},
repeat: None,
depth: ctx.depth + 1,
};
self.stack.push((type_name, ctor_name));
let block = func.block.clone();
self.walk_block(&block, &mut child_env, &child_ctx);
self.stack.pop();
true
}
fn resolve_vb_arg(
&mut self,
args: &syn::punctuated::Punctuated<syn::Expr, syn::Token![,]>,
index: usize,
env: &Env,
ctx: &Ctx,
) -> Option<VbVal> {
args.iter()
.nth(index)
.and_then(|arg| self.eval_vb(arg, env, ctx))
}
fn eval_vb(&mut self, expr: &syn::Expr, env: &Env, ctx: &Ctx) -> Option<VbVal> {
let expr = unwrap_expr(expr);
match expr {
syn::Expr::Path(p) => {
let name = p.path.segments.last()?.ident.to_string();
env.vb.get(&name).cloned()
}
syn::Expr::MethodCall(mc) => {
let method = mc.method.to_string();
if method == "then" || method == "then_some" {
let arg = mc.args.first()?;
let inner = match unwrap_expr(arg) {
syn::Expr::Closure(c) => c.body.as_ref(),
other => other,
};
let val = self.eval_vb(inner, env, ctx)?;
let reason = format!("only when {}", crate::load::type_text(&mc.receiver));
return Some(VbVal {
certainty: merge(val.certainty, Certainty::Conditional(reason)),
..val
});
}
if method == "map" || method == "and_then" {
let base = self.eval_vb(&mc.receiver, env, ctx)?;
let arg = mc.args.first()?;
let syn::Expr::Closure(closure) = unwrap_expr(arg) else {
return None;
};
let param = closure.inputs.first().and_then(binding_name)?;
let mut scoped = env.clone();
scoped.vb.insert(param, base.clone());
let val = self.eval_vb(&closure.body, &scoped, ctx)?;
return Some(VbVal {
certainty: merge(val.certainty, base.certainty),
..val
});
}
if matches!(
method.as_str(),
"as_ref" | "as_mut" | "clone" | "to_owned" | "to_dtype"
) {
return self.eval_vb(&mc.receiver, env, ctx);
}
let op = known::prefix_method(&method)?;
let base = self.eval_vb(&mc.receiver, env, ctx)?;
match op {
PrefixOp::Root => Some(VbVal {
key: Key::default(),
..base
}),
PrefixOp::Push | PrefixOp::Replace => {
let arg = mc.args.first()?;
let (segs, seg_certainty) = self.eval_prefix_arg(arg, ctx);
let start = if op == PrefixOp::Replace {
Key::default()
} else {
base.key
};
Some(VbVal {
root: base.root,
key: start.extend(&segs),
certainty: merge(base.certainty, seg_certainty),
})
}
}
}
syn::Expr::Call(c) => {
let path = call_path(&c.func)?;
if path.last().map(String::as_str) == Some("Some") && c.args.len() == 1 {
return self.eval_vb(&c.args[0], env, ctx);
}
None
}
_ => None,
}
}
fn eval_prefix_arg(&mut self, expr: &syn::Expr, ctx: &Ctx) -> (Vec<KeySeg>, Certainty) {
let expr = unwrap_expr(expr);
if let Some(text) = string_literal(expr) {
return (literal_segs(&text), Certainty::Certain);
}
if let syn::Expr::Macro(m) = expr {
if m.mac
.path
.segments
.last()
.map(|s| s.ident.to_string())
.as_deref()
== Some("format")
{
if let Some(segs) = format_segs(&m.mac.tokens.to_string()) {
return (segs, Certainty::Certain);
}
}
}
let text = crate::load::type_text(expr);
self.out.diagnose(
span_of_expr(ctx, expr),
format!("prefix argument `{text}` is not a literal or `format!`; segment is dynamic"),
None,
);
(vec![KeySeg::Dynamic { expr: text }], Certainty::Certain)
}
}
fn combine(a: &Certainty, b: &Certainty, unconditional: bool, reason: &str) -> Certainty {
let merged = merge(a.clone(), b.clone());
if unconditional {
merged
} else {
match merged {
Certainty::Certain => Certainty::Conditional(reason.to_string()),
other => other,
}
}
}
fn merge(a: Certainty, b: Certainty) -> Certainty {
match (&a, &b) {
(Certainty::Unknown(_), _) => a,
(_, Certainty::Unknown(_)) => b,
(Certainty::Conditional(_), _) => a,
(_, Certainty::Conditional(_)) => b,
_ => Certainty::Certain,
}
}
fn literal_segs(text: &str) -> Vec<KeySeg> {
text.split('.')
.filter(|p| !p.is_empty())
.map(|p| KeySeg::Literal(p.to_string()))
.collect()
}
fn format_segs(tokens: &str) -> Option<Vec<KeySeg>> {
let tokens = tokens.trim();
let rest = tokens.strip_prefix('"')?;
let end = find_unescaped_quote(rest)?;
let template = &rest[..end];
let args: Vec<String> = rest[end + 1..]
.trim_start()
.trim_start_matches(',')
.split(',')
.map(|a| a.trim().to_string())
.filter(|a| !a.is_empty())
.collect();
let mut positional = args.into_iter();
let mut segs = Vec::new();
for part in template.split('.').filter(|p| !p.is_empty()) {
if let Some(inner) = brace_content(part) {
let expr = if inner.is_empty() {
positional.next().unwrap_or_else(|| "_".to_string())
} else {
inner.split(':').next().unwrap_or(inner).to_string()
};
segs.push(KeySeg::Dynamic { expr });
} else if part.contains('{') {
segs.push(KeySeg::Template {
text: part.to_string(),
});
} else {
segs.push(KeySeg::Literal(part.to_string()));
}
}
Some(segs)
}
fn brace_content(part: &str) -> Option<&str> {
let inner = part.strip_prefix('{')?.strip_suffix('}')?;
(!inner.contains('{')).then_some(inner)
}
fn find_unescaped_quote(s: &str) -> Option<usize> {
let bytes = s.as_bytes();
let mut i = 0;
while i < bytes.len() {
match bytes[i] {
b'\\' => i += 2,
b'"' => return Some(i),
_ => i += 1,
}
}
None
}
fn macro_cannot_register_params(path: &syn::Path) -> bool {
let Some(name) = path
.segments
.last()
.map(|segment| segment.ident.to_string())
else {
return false;
};
matches!(
name.as_str(),
"assert"
| "assert_eq"
| "assert_ne"
| "debug_assert"
| "debug_assert_eq"
| "debug_assert_ne"
| "bail"
| "ensure"
| "panic"
| "todo"
| "unimplemented"
| "unreachable"
| "print"
| "println"
| "eprint"
| "eprintln"
| "dbg"
)
}
fn string_literal(expr: &syn::Expr) -> Option<String> {
match unwrap_expr(expr) {
syn::Expr::Lit(lit) => match &lit.lit {
syn::Lit::Str(s) => Some(s.value()),
_ => None,
},
_ => None,
}
}
fn unwrap_expr(expr: &syn::Expr) -> &syn::Expr {
match expr {
syn::Expr::Try(t) => unwrap_expr(&t.expr),
syn::Expr::Reference(r) => unwrap_expr(&r.expr),
syn::Expr::Paren(p) => unwrap_expr(&p.expr),
syn::Expr::Group(g) => unwrap_expr(&g.expr),
other => other,
}
}
fn call_path(expr: &syn::Expr) -> Option<Vec<String>> {
match unwrap_expr(expr) {
syn::Expr::Path(p) => Some(
p.path
.segments
.iter()
.map(|s| s.ident.to_string())
.collect(),
),
_ => None,
}
}
fn known_constructor(path: &[String], func: &str) -> Option<&'static crate::known::Constructor> {
let namespaced = path.len() >= 2
&& matches!(
path.first().map(String::as_str),
Some("candle_nn" | "nn" | "candle")
);
namespaced.then(|| known::lookup(func)).flatten()
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
struct RnnConfig {
biases: Option<bool>,
default_names: bool,
}
fn rnn_config(
func: &str,
args: &syn::punctuated::Punctuated<syn::Expr, syn::Token![,]>,
) -> RnnConfig {
if !matches!(func, "lstm" | "gru") {
return RnnConfig::default();
}
let Some(text) = args.iter().nth(2).map(crate::load::type_text) else {
return RnnConfig::default();
};
let text = text.replace(' ', "");
match text.rsplit("::").next().unwrap_or(&text) {
"default()" => RnnConfig {
biases: Some(true),
default_names: true,
},
"default_no_bias()" => RnnConfig {
biases: Some(false),
default_names: true,
},
_ => RnnConfig::default(),
}
}
fn config_named_family(leaf: &str) -> String {
match leaf.strip_suffix("l0") {
Some(stem) => format!("{stem}l{{layer_idx}}{{direction}}"),
None => leaf.to_string(),
}
}
fn constructor_leaf_shape(
func: &str,
leaf: &str,
args: &syn::punctuated::Punctuated<syn::Expr, syn::Token![,]>,
) -> Option<String> {
let arg = |index: usize| args.iter().nth(index).map(crate::load::type_text);
let pair =
|left: Option<String>, right: Option<String>| Some(format!("({}, {})", left?, right?));
match (func, leaf) {
("linear" | "linear_no_bias" | "linear_b", "weight") => pair(arg(1), arg(0)),
("linear" | "linear_b", "bias") => arg(1),
("embedding", "weight") => pair(arg(0), arg(1)),
(
"layer_norm" | "layer_norm_no_bias" | "rms_norm" | "batch_norm" | "group_norm",
"weight" | "bias" | "running_mean" | "running_var",
) => {
if func == "group_norm" {
arg(1)
} else {
arg(0)
}
}
("prelu", "weight") => Some(format!("({}.unwrap_or(1),)", arg(0)?)),
("lstm" | "gru", leaf) if leaf.starts_with("weight_") || leaf.starts_with("bias_") => {
let gates = if func == "lstm" { 4 } else { 3 };
let rows = format!("{gates} * {}", arg(1)?);
match (leaf.starts_with("weight_"), leaf.contains("_ih_")) {
(true, true) => Some(format!("({rows}, {})", arg(0)?)),
(true, false) => Some(format!("({rows}, {})", arg(1)?)),
(false, _) => Some(format!("({rows},)")),
}
}
("conv1d" | "conv1d_no_bias", "weight") => Some(format!(
"({}, {} / groups({}), {})",
arg(1)?,
arg(0)?,
arg(3)?,
arg(2)?
)),
("conv2d" | "conv2d_no_bias", "weight") => Some(format!(
"({}, {} / groups({}), {}, {})",
arg(1)?,
arg(0)?,
arg(3)?,
arg(2)?,
arg(2)?
)),
("conv_transpose1d" | "conv_transpose1d_no_bias", "weight") => Some(format!(
"({}, {} / groups({}), {})",
arg(0)?,
arg(1)?,
arg(3)?,
arg(2)?
)),
("conv_transpose2d" | "conv_transpose2d_no_bias", "weight") => Some(format!(
"({}, {}, {}, {})",
arg(0)?,
arg(1)?,
arg(2)?,
arg(2)?
)),
(
"conv1d"
| "conv1d_no_bias"
| "conv2d"
| "conv2d_no_bias"
| "conv_transpose1d"
| "conv_transpose1d_no_bias"
| "conv_transpose2d"
| "conv_transpose2d_no_bias",
"bias",
) => arg(1),
_ => None,
}
}
fn normalize_qualified_path(path: &[String]) -> String {
let mut path = path;
while matches!(path.first().map(String::as_str), Some("crate" | "self")) {
path = &path[1..];
}
path.join("::")
}
fn normalize_qualified_name(name: &str) -> String {
normalize_qualified_path(
&name
.split("::")
.map(ToString::to_string)
.collect::<Vec<_>>(),
)
}
fn looks_like_iteration(expr: &syn::Expr) -> bool {
match unwrap_expr(expr) {
syn::Expr::Range(_) => true,
syn::Expr::MethodCall(mc) => matches!(
mc.method.to_string().as_str(),
"iter"
| "into_iter"
| "iter_mut"
| "enumerate"
| "zip"
| "take"
| "skip"
| "rev"
| "filter"
| "chain"
| "step_by"
| "windows"
| "chunks"
),
_ => false,
}
}
fn binding_name(pat: &syn::Pat) -> Option<String> {
match pat {
syn::Pat::Ident(id) => Some(id.ident.to_string()),
syn::Pat::Type(t) => binding_name(&t.pat),
_ => None,
}
}
fn span_of_expr<T: syn::spanned::Spanned>(ctx: &Ctx, node: &T) -> SrcSpan {
let start = node.span().start();
SrcSpan {
file: ctx.file,
line: start.line,
col: start.column,
}
}