Skip to content
32 changes: 7 additions & 25 deletions compiler/rustc_builtin_macros/src/asm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -588,13 +588,7 @@ pub(super) fn expand_asm<'cx>(
return ExpandResult::Retry(());
};
let expr = match mac {
Ok(inline_asm) => Box::new(ast::Expr {
id: ast::DUMMY_NODE_ID,
kind: ast::ExprKind::InlineAsm(Box::new(inline_asm)),
span: sp,
attrs: ast::AttrVec::new(),
tokens: None,
}),
Ok(inline_asm) => ecx.expr(sp, ast::ExprKind::InlineAsm(Box::new(inline_asm))),
Err(guar) => DummyResult::raw_expr(sp, Some(guar)),
};
MacEager::expr(expr)
Expand All @@ -618,13 +612,7 @@ pub(super) fn expand_naked_asm<'cx>(
return ExpandResult::Retry(());
};
let expr = match mac {
Ok(inline_asm) => Box::new(ast::Expr {
id: ast::DUMMY_NODE_ID,
kind: ast::ExprKind::InlineAsm(Box::new(inline_asm)),
span: sp,
attrs: ast::AttrVec::new(),
tokens: None,
}),
Ok(inline_asm) => ecx.expr(sp, ast::ExprKind::InlineAsm(Box::new(inline_asm))),
Err(guar) => DummyResult::raw_expr(sp, Some(guar)),
};
MacEager::expr(expr)
Expand All @@ -648,17 +636,11 @@ pub(super) fn expand_global_asm<'cx>(
return ExpandResult::Retry(());
};
match mac {
Ok(inline_asm) => MacEager::items(smallvec![Box::new(ast::Item {
attrs: ast::AttrVec::new(),
id: ast::DUMMY_NODE_ID,
kind: ast::ItemKind::GlobalAsm(Box::new(inline_asm)),
vis: ast::Visibility {
span: sp.shrink_to_lo(),
kind: ast::VisibilityKind::Inherited,
},
span: sp,
tokens: None,
})]),
Ok(inline_asm) => MacEager::items(smallvec![ecx.item(
sp,
ast::AttrVec::new(),
ast::ItemKind::GlobalAsm(Box::new(inline_asm))
)]),
Err(guar) => DummyResult::any(sp, guar),
}
}
Expand Down
24 changes: 5 additions & 19 deletions compiler/rustc_builtin_macros/src/assert.rs
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
mod context;

use rustc_ast::token::Delimiter;
use rustc_ast::tokenstream::{DelimSpan, TokenStream};
use rustc_ast::{DelimArgs, Expr, ExprKind, MacCall, Path, PathSegment, UnOp, token};
use rustc_ast::tokenstream::TokenStream;
use rustc_ast::{Expr, ExprKind, Path, UnOp, token};
use rustc_ast_pretty::pprust;
use rustc_errors::PResult;
use rustc_expand::base::{DummyResult, ExpandResult, ExtCtxt, MacEager, MacroExpanderResult};
Expand Down Expand Up @@ -34,14 +34,7 @@ pub(crate) fn expand_assert<'cx>(
let panic_path = || {
if use_panic_2021(span) {
// On edition 2021, we always call `$crate::panic::panic_2021!()`.
Path {
span: call_site_span,
segments: cx
.std_path(&[sym::panic, sym::panic_2021])
.into_iter()
.map(PathSegment::from_ident)
.collect(),
}
cx.path(call_site_span, cx.std_path(&[sym::panic, sym::panic_2021]))
} else {
// Before edition 2021, we call `panic!()` unqualified,
// such that it calls either `std::panic!()` or `core::panic!()`.
Expand All @@ -51,16 +44,9 @@ pub(crate) fn expand_assert<'cx>(

// Simply uses the user provided message instead of generating custom outputs
let expr = if let Some(tokens) = custom_message {
let then = cx.expr(
let then = cx.expr_macro_call(
call_site_span,
ExprKind::MacCall(Box::new(MacCall {
path: panic_path(),
args: Box::new(DelimArgs {
dspan: DelimSpan::from_single(call_site_span),
delim: Delimiter::Parenthesis,
tokens,
}),
})),
cx.macro_call(call_site_span, panic_path(), Delimiter::Parenthesis, tokens),
);
expr_if_not(cx, call_site_span, cond_expr, then, None)
}
Expand Down
31 changes: 6 additions & 25 deletions compiler/rustc_builtin_macros/src/assert/context.rs
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
use rustc_ast::token::{self, Delimiter, IdentKind};
use rustc_ast::tokenstream::{DelimSpan, TokenStream, TokenTree};
use rustc_ast::{
BinOpKind, BorrowKind, DUMMY_NODE_ID, DelimArgs, Expr, ExprKind, ItemKind, MacCall, MethodCall,
Mutability, Path, PathSegment, Stmt, StructRest, UnOp, UseTree, UseTreeAndId, UseTreeKind,
BinOpKind, BorrowKind, DUMMY_NODE_ID, DelimArgs, Expr, ExprKind, ItemKind, MacCall, Mutability,
Path, Stmt, StructRest, UnOp, UseTree, UseTreeAndId, UseTreeKind,
};
use rustc_ast_pretty::pprust;
use rustc_data_structures::fx::FxHashSet;
Expand Down Expand Up @@ -382,20 +382,15 @@ impl<'cx, 'a> Context<'cx, 'a> {
);
let try_capture_call = self
.cx
.stmt_expr(expr_method_call(
self.cx,
PathSegment {
args: None,
id: DUMMY_NODE_ID,
ident: Ident::new(sym::try_capture, self.span),
},
expr_paren(self.cx, self.span, self.cx.expr_addr_of(self.span, wrapper)),
.stmt_expr(self.cx.expr_method_call(
self.span,
self.cx.expr_paren(self.span, self.cx.expr_addr_of(self.span, wrapper)),
Ident::new(sym::try_capture, self.span),
thin_vec![expr_addr_of_mut(
self.cx,
self.span,
self.cx.expr_path(Path::from_ident(capture)),
)],
self.span,
))
.add_trailing_semicolon();
let local_bind_path = self.cx.expr_path(Path::from_ident(local_bind));
Expand Down Expand Up @@ -448,17 +443,3 @@ fn escape_to_fmt(s: &str) -> String {
fn expr_addr_of_mut(cx: &ExtCtxt<'_>, sp: Span, e: Box<Expr>) -> Box<Expr> {
cx.expr(sp, ExprKind::AddrOf(BorrowKind::Ref, Mutability::Mut, e))
}

fn expr_method_call(
cx: &ExtCtxt<'_>,
seg: PathSegment,
receiver: Box<Expr>,
args: ThinVec<Box<Expr>>,
span: Span,
) -> Box<Expr> {
cx.expr(span, ExprKind::MethodCall(Box::new(MethodCall { seg, receiver, args, span })))
}

fn expr_paren(cx: &ExtCtxt<'_>, sp: Span, e: Box<Expr>) -> Box<Expr> {
cx.expr(sp, ExprKind::Paren(e))
}
108 changes: 29 additions & 79 deletions compiler/rustc_builtin_macros/src/autodiff.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,8 @@ mod llvm_enzyme {
use rustc_ast::tokenstream::*;
use rustc_ast::visit::AssocCtxt::*;
use rustc_ast::{
self as ast, AngleBracketedArg, AngleBracketedArgs, AnonConst, AssocItemKind, BindingMode,
FnRetTy, FnSig, GenericArg, GenericArgs, GenericParamKind, Generics, ItemKind,
MetaItemInner, PatKind, Path, PathSegment, TyKind, Visibility,
self as ast, AnonConst, FnRetTy, FnSig, GenericArg, GenericParamKind, Generics, ItemKind,
MetaItemInner, PatKind, TyKind, Visibility,
};
use rustc_attr_ir::RustcAutodiff;
use rustc_expand::base::{Annotatable, ExtCtxt};
Expand Down Expand Up @@ -154,7 +153,7 @@ mod llvm_enzyme {
}

fn meta_item_inner_to_ts(t: &MetaItemInner, ts: &mut Vec<TokenTree>) {
let comma: Token = Token::new(TokenKind::Comma, Span::default());
let comma = Token::new(TokenKind::Comma, Span::default());
let val = first_ident(t);
let t = Token::from_ast_ident(val);
ts.push(TokenTree::Token(t, Spacing::Joint));
Expand Down Expand Up @@ -275,7 +274,7 @@ mod llvm_enzyme {
// Now, if the user gave a width (vector aka batch-mode ad), then we copy it.
// If it is not given, we default to 1 (scalar mode).
let start_position;
let kind: LitKind = LitKind::Integer;
let kind = LitKind::Integer;
let symbol;
if meta_item_vec.len() >= 2
&& let Some(width) = width(&meta_item_vec[1])
Expand All @@ -287,7 +286,7 @@ mod llvm_enzyme {
symbol = sym::integer(1);
}

let l: Lit = Lit { kind, symbol, suffix: None };
let l = Lit { kind, symbol, suffix: None };
let t = Token::new(TokenKind::Literal(l), Span::default());
let comma = Token::new(TokenKind::Comma, Span::default());
ts.push(TokenTree::Token(t, Spacing::Joint));
Expand All @@ -306,9 +305,9 @@ mod llvm_enzyme {
}
// We remove the last, trailing comma.
ts.pop();
let ts: TokenStream = TokenStream::from_iter(ts);
let ts = TokenStream::from_iter(ts);

let x: RustcAutodiff = from_ast(ecx, &meta_item_vec, has_ret, mode);
let x = from_ast(ecx, &meta_item_vec, has_ret, mode);
if !x.is_active() {
// We encountered an error, so we return the original item.
// This allows us to potentially parse other attributes.
Expand Down Expand Up @@ -346,7 +345,7 @@ mod llvm_enzyme {
let mut rustc_ad_attr =
Box::new(ast::NormalAttr::from_ident(Ident::with_dummy_span(sym::rustc_autodiff)));

let ts2: Vec<TokenTree> = vec![TokenTree::Token(
let ts2 = vec![TokenTree::Token(
Token::new(TokenKind::Ident(sym::never, IdentKind::Normal), span),
Spacing::Joint,
)];
Expand Down Expand Up @@ -382,7 +381,7 @@ mod llvm_enzyme {
let mut has_inline_never = false;

// Don't add it multiple times:
let orig_annotatable: Annotatable = match item {
let orig_annotatable = match item {
Annotatable::Item(ref mut iitem) => {
if !iitem.attrs.iter().any(|a| same_attribute(&a.kind, &attr.kind)) {
iitem.attrs.push(attr);
Expand Down Expand Up @@ -439,7 +438,7 @@ mod llvm_enzyme {

let d_annotatable = match &item {
Annotatable::AssocItem(_, ctxt) => {
let assoc_item: AssocItemKind = ast::AssocItemKind::Fn(d_fn);
let assoc_item = ast::AssocItemKind::Fn(d_fn);
let d_fn = Box::new(ast::AssocItem {
attrs: d_attrs,
id: ast::DUMMY_NODE_ID,
Expand All @@ -459,12 +458,7 @@ mod llvm_enzyme {
Annotatable::Stmt(_) => {
let mut d_fn = ecx.item(span, d_attrs, ItemKind::Fn(d_fn));
d_fn.vis = vis;

Annotatable::Stmt(Box::new(ast::Stmt {
id: ast::DUMMY_NODE_ID,
kind: ast::StmtKind::Item(d_fn),
span,
}))
Annotatable::Stmt(Box::new(ecx.stmt_item(span, d_fn)))
}
_ => {
unreachable!("item kind checked previously")
Expand Down Expand Up @@ -568,7 +562,7 @@ mod llvm_enzyme {
let call_expr = ecx.expr_call(
span,
ecx.expr_path(enzyme_path),
vec![primal_fn_ptr, diff_path_expr, tuple_expr].into(),
thin_vec![primal_fn_ptr, diff_path_expr, tuple_expr],
);

ecx.stmt_expr(call_expr)
Expand All @@ -591,35 +585,20 @@ mod llvm_enzyme {
GenericParamKind::Type { .. } => {
let path = ast::Path::from_ident(p.ident);
let ty = ecx.ty_path(path);
Some(AngleBracketedArg::Arg(GenericArg::Type(ty)))
Some(GenericArg::Type(ty))
}
GenericParamKind::Const { .. } => {
let expr = ecx.expr_path(ast::Path::from_ident(p.ident));
let anon_const = AnonConst { id: ast::DUMMY_NODE_ID, value: expr };
Some(AngleBracketedArg::Arg(GenericArg::Const(anon_const)))
Some(GenericArg::Const(anon_const))
}
GenericParamKind::Lifetime => None,
})
.collect::<ThinVec<_>>();

let args: AngleBracketedArgs = AngleBracketedArgs { span, args: generic_args };

let segment = PathSegment {
ident,
id: ast::DUMMY_NODE_ID,
args: Some(Box::new(GenericArgs::AngleBracketed(args))),
};

let segments = if is_impl {
thin_vec![
PathSegment { ident: Ident::from_str("Self"), id: ast::DUMMY_NODE_ID, args: None },
segment,
]
} else {
thin_vec![segment]
};
.collect::<Vec<_>>();

let path = Path { span, segments };
let idents =
if is_impl { vec![Ident::new(kw::SelfUpper, span), ident] } else { vec![ident] };
let path = ecx.path_all(span, false, idents, generic_args);

ecx.expr_path(path)
}
Expand Down Expand Up @@ -657,9 +636,7 @@ mod llvm_enzyme {
assert!(sig.decl.inputs.len() == x.input_activity.len());
assert!(has_ret == x.has_ret_activity());
let mut d_decl = sig.decl.clone();
let mut d_inputs = Vec::new();
let mut new_inputs = Vec::new();
let mut idents = Vec::new();
let mut d_inputs = ThinVec::new();
let mut act_ret = ThinVec::new();

// We have two loops, a first one just to check the activities and types and possibly report
Expand Down Expand Up @@ -724,15 +701,10 @@ mod llvm_enzyme {
debug!("{:#?}", &shadow_arg.pat);
panic!("not an ident?");
};
let name: String = format!("d{}_{}", old_name, i);
new_inputs.push(name.clone());
let name = format!("d{}_{}", old_name, i);
let ident = Ident::from_str_and_span(&name, shadow_arg.pat.span);
*shadow_arg.pat = ast::Pat {
id: ast::DUMMY_NODE_ID,
kind: PatKind::Ident(BindingMode::NONE, ident, None),
span: shadow_arg.pat.span,
};
d_inputs.push(shadow_arg.clone());
*shadow_arg.pat = ecx.pat_ident(shadow_arg.pat.span, ident);
d_inputs.push(shadow_arg);
}
}
DiffActivity::Dual
Expand All @@ -755,15 +727,11 @@ mod llvm_enzyme {
debug!("{:#?}", &shadow_arg.pat);
panic!("not an ident?");
};
let name: String = format!("b{}_{}", old_name, i);
new_inputs.push(name.clone());
let name = format!("b{}_{}", old_name, i);
let ident = Ident::from_str_and_span(&name, shadow_arg.pat.span);
*shadow_arg.pat = ast::Pat {
id: ast::DUMMY_NODE_ID,
kind: PatKind::Ident(BindingMode::NONE, ident, None),
span: shadow_arg.pat.span,
};
d_inputs.push(shadow_arg.clone());

*shadow_arg.pat = ecx.pat_ident(shadow_arg.pat.span, ident);
d_inputs.push(shadow_arg);
}
}
DiffActivity::Const => {
Expand All @@ -773,11 +741,6 @@ mod llvm_enzyme {
panic!("Should not happen");
}
}
if let PatKind::Ident(_, ident, _) = arg.pat.kind {
idents.push(ident);
} else {
panic!("not an ident?");
}
}

let active_only_ret = x.ret_activity == DiffActivity::ActiveOnly;
Expand All @@ -798,33 +761,20 @@ mod llvm_enzyme {
};
let name = "dret".to_string();
let ident = Ident::from_str_and_span(&name, ty.span);
let shadow_arg = ast::Param {
attrs: ThinVec::new(),
ty: ty.clone(),
pat: Box::new(ast::Pat {
id: ast::DUMMY_NODE_ID,
kind: PatKind::Ident(BindingMode::NONE, ident, None),
span: ty.span,
}),
id: ast::DUMMY_NODE_ID,
span: ty.span,
is_placeholder: false,
};
let shadow_arg = ecx.param(ty.span, ident, ty);
d_inputs.push(shadow_arg);
new_inputs.push(name);
}
_ => {}
}
}
d_decl.inputs = d_inputs.into();
d_decl.inputs = d_inputs;

if x.mode.is_fwd() {
let ty = match d_decl.output {
FnRetTy::Ty(ref ty) => ty.clone(),
FnRetTy::Default(span) => {
// We want to return std::hint::black_box(()).
let kind = TyKind::Tup(ThinVec::new());
let ty = Box::new(rustc_ast::Ty { kind, id: ast::DUMMY_NODE_ID, span });
let ty = ecx.ty_unit(span);
d_decl.output = FnRetTy::Ty(ty.clone());
assert!(matches!(x.ret_activity, DiffActivity::None));
// this won't be used below, so any type would be fine.
Expand Down
Loading
Loading