pub use self::{binary::*, unary::*};
pub(crate) mod binary;
pub(crate) mod unary;
use crate::ops::Methods;
use core::str::FromStr;
use proc_macro2::TokenStream;
use quote::quote;
use syn::parse;
use syn::punctuated::Punctuated;
use syn::token::Comma;
use syn::{Expr, ExprCall, Ident};
pub fn handle_expr(expr: &Expr, variable: &Ident) -> TokenStream {
match expr {
Expr::Array(inner) => {
let grad = inner
.elems
.iter()
.map(|e| parse::<Expr>(handle_expr(e, variable).into()).unwrap());
quote! {
[#(#grad),*]
}
}
Expr::Binary(inner) => handle_binary(inner, variable),
Expr::Call(inner) => handle_call(inner, variable),
Expr::Closure(inner) => handle_expr(&inner.body, variable),
Expr::Const(_) => {
quote! { T::default() }
}
Expr::Group(inner) => handle_expr(&inner.expr, variable),
Expr::Lit(inner) => match &inner.lit {
syn::Lit::Str(literal) => {
match parse::<syn::ItemFn>(TokenStream::from_str(&literal.value()).unwrap().into())
{
Ok(item) => crate::handle::item::handle_item_fn(&item, variable),
Err(_) => quote! { 0.0 },
}
}
_ => quote! { 0.0 },
},
Expr::MethodCall(inner) => Methods::from_method_call(inner, variable),
Expr::Paren(inner) => handle_expr(&inner.expr, variable),
Expr::Path(inner) => {
let syn::ExprPath { path, .. } = inner;
if path.segments.len() != 1 {
panic!("Unsupported path!");
}
if path.segments[0].ident == *variable {
quote! { 1.0 }
} else {
quote! { 0.0 }
}
}
Expr::Reference(inner) => handle_expr(&inner.expr, variable),
Expr::Unary(inner) => handle_unary(inner, variable),
_ => panic!("Unsupported expression!"),
}
}
pub fn handle_call(expr: &ExprCall, var: &Ident) -> TokenStream {
let ExprCall { args, func, .. } = expr;
let mut grad = quote! { 0.0 };
for arg in args {
let arg = handle_expr(arg, var);
grad = quote! { #grad + #arg };
}
let df = handle_expr(func, var);
quote! { #df + #grad }
}
#[allow(dead_code)]
fn grad_ctx_with_args(ctx: &Box<Expr>, args: &Punctuated<Expr, Comma>, var: &Ident) -> TokenStream {
let grad = handle_expr(ctx, var);
let da = punctuated_grad(args, var);
quote! { #grad + #da }
}
fn punctuated_grad(args: &Punctuated<Expr, Comma>, var: &Ident) -> TokenStream {
args.iter()
.map(|arg| handle_expr(arg, var))
.fold(quote! { 0.0 }, |acc, arg| quote! { #acc + #arg })
}