1mod llvm_enzyme {
7 use std::str::FromStr;
8 use std::string::String;
9
10 use rustc_ast::expand::autodiff_attrs::{
11 DiffActivity, DiffMode, valid_input_activity, valid_ret_activity, valid_ty_for_activity,
12 };
13 use rustc_ast::token::{IdentKind, Lit, LitKind, Token, TokenKind};
14 use rustc_ast::tokenstream::*;
15 use rustc_ast::visit::AssocCtxt::*;
16 use rustc_ast::{
17 self as ast, AngleBracketedArg, AngleBracketedArgs, AnonConst, AssocItemKind, BindingMode,
18 FnRetTy, FnSig, GenericArg, GenericArgs, GenericParamKind, Generics, ItemKind,
19 MetaItemInner, PatKind, Path, PathSegment, TyKind, Visibility,
20 };
21 use rustc_attr_ir::RustcAutodiff;
22 use rustc_expand::base::{Annotatable, ExtCtxt};
23 use rustc_span::{DUMMY_SP, Ident, Span, Symbol, kw, sym};
24 use thin_vec::{ThinVec, thin_vec};
25 use tracing::{debug, trace};
26
27 use crate::diagnostics;
28
29 pub(crate) fn outer_normal_attr(
30 kind: &Box<rustc_ast::NormalAttr>,
31 id: rustc_ast::AttrId,
32 span: Span,
33 ) -> rustc_ast::Attribute {
34 let style = rustc_ast::AttrStyle::Outer;
35 let kind = rustc_ast::AttrKind::Normal(kind.clone());
36 rustc_ast::Attribute { kind, id, style, span }
37 }
38
39 fn has_ret(ty: &FnRetTy) -> bool {
42 match ty {
43 FnRetTy::Ty(ty) => !ty.kind.is_unit(),
44 FnRetTy::Default(_) => false,
45 }
46 }
47 fn first_ident(x: &MetaItemInner) -> rustc_span::Ident {
48 if let Some(l) = x.lit() {
49 match l.kind {
50 ast::LitKind::Int(val, _) => {
51 return rustc_span::Ident::from_str(val.get().to_string().as_str());
53 }
54 _ => {}
55 }
56 }
57
58 let segments = &x.meta_item().unwrap().path.segments;
59 if !(segments.len() == 1) {
::core::panicking::panic("assertion failed: segments.len() == 1")
};assert!(segments.len() == 1);
60 segments[0].ident
61 }
62
63 fn name(x: &MetaItemInner) -> String {
64 first_ident(x).name.to_string()
65 }
66
67 fn width(x: &MetaItemInner) -> Option<u128> {
68 let lit = x.lit()?;
69 match lit.kind {
70 ast::LitKind::Int(x, _) => Some(x.get()),
71 _ => None,
72 }
73 }
74
75 fn extract_item_info(iitem: &Box<ast::Item>) -> Option<(Visibility, FnSig, Ident, Generics)> {
77 match &iitem.kind {
78 ItemKind::Fn(ast::Fn { sig, ident, generics, .. }) => {
79 Some((iitem.vis.clone(), sig.clone(), *ident, generics.clone()))
80 }
81 _ => None,
82 }
83 }
84
85 pub(crate) fn from_ast(
86 ecx: &mut ExtCtxt<'_>,
87 meta_item: &ThinVec<MetaItemInner>,
88 has_ret: bool,
89 mode: DiffMode,
90 ) -> RustcAutodiff {
91 let dcx = ecx.sess.dcx();
92
93 let mut first_activity = 1;
96
97 let width = if let [_, x, ..] = &meta_item[..]
98 && let Some(x) = width(x)
99 {
100 first_activity = 2;
101 match x.try_into() {
102 Ok(x) => x,
103 Err(_) => {
104 dcx.emit_err(diagnostics::AutoDiffInvalidWidth {
105 span: meta_item[1].span(),
106 width: x,
107 });
108 return RustcAutodiff::error();
109 }
110 }
111 } else {
112 1
113 };
114
115 let mut activities: Vec<DiffActivity> = ::alloc::vec::Vec::new()vec![];
116 let mut errors = false;
117 for x in &meta_item[first_activity..] {
118 let activity_str = name(x);
119 let res = DiffActivity::from_str(&activity_str);
120 match res {
121 Ok(x) => activities.push(x),
122 Err(_) => {
123 dcx.emit_err(diagnostics::AutoDiffUnknownActivity {
124 span: x.span(),
125 act: activity_str,
126 });
127 errors = true;
128 }
129 };
130 }
131 if errors {
132 return RustcAutodiff::error();
133 }
134
135 let (ret_activity, input_activity) = if has_ret {
138 let Some((last, rest)) = activities.split_last() else {
139 {
::core::panicking::panic_fmt(format_args!("internal error: entered unreachable code: {0}",
format_args!("should not be reachable because we counted the number of activities previously")));
};unreachable!(
140 "should not be reachable because we counted the number of activities previously"
141 );
142 };
143 (last, rest)
144 } else {
145 (&DiffActivity::None, activities.as_slice())
146 };
147
148 RustcAutodiff {
149 mode,
150 width,
151 ret_activity: *ret_activity,
152 input_activity: input_activity.iter().cloned().collect(),
153 }
154 }
155
156 fn meta_item_inner_to_ts(t: &MetaItemInner, ts: &mut Vec<TokenTree>) {
157 let comma: Token = Token::new(TokenKind::Comma, Span::default());
158 let val = first_ident(t);
159 let t = Token::from_ast_ident(val);
160 ts.push(TokenTree::Token(t, Spacing::Joint));
161 ts.push(TokenTree::Token(comma, Spacing::Alone));
162 }
163
164 pub(crate) fn expand_forward(
165 ecx: &mut ExtCtxt<'_>,
166 expand_span: Span,
167 meta_item: &ast::MetaItem,
168 item: Annotatable,
169 ) -> Vec<Annotatable> {
170 expand_with_mode(ecx, expand_span, meta_item, item, DiffMode::Forward)
171 }
172
173 pub(crate) fn expand_reverse(
174 ecx: &mut ExtCtxt<'_>,
175 expand_span: Span,
176 meta_item: &ast::MetaItem,
177 item: Annotatable,
178 ) -> Vec<Annotatable> {
179 expand_with_mode(ecx, expand_span, meta_item, item, DiffMode::Reverse)
180 }
181
182 pub(crate) fn expand_with_mode(
206 ecx: &mut ExtCtxt<'_>,
207 expand_span: Span,
208 meta_item: &ast::MetaItem,
209 mut item: Annotatable,
210 mode: DiffMode,
211 ) -> Vec<Annotatable> {
212 let dcx = ecx.sess.dcx();
213
214 let Some((vis, sig, primal, generics, is_impl)) = (match &item {
218 Annotatable::Item(iitem) => {
219 extract_item_info(iitem).map(|(v, s, p, g)| (v, s, p, g, false))
220 }
221 Annotatable::Stmt(stmt) => match &stmt.kind {
222 ast::StmtKind::Item(iitem) => {
223 extract_item_info(iitem).map(|(v, s, p, g)| (v, s, p, g, false))
224 }
225 _ => None,
226 },
227 Annotatable::AssocItem(assoc_item, _ctxt @ (Impl { of_trait: _ } | Trait)) => {
228 match &assoc_item.kind {
229 ast::AssocItemKind::Fn(ast::Fn { sig, ident, generics, .. }) => {
230 Some((assoc_item.vis.clone(), sig.clone(), *ident, generics.clone(), true))
231 }
232 _ => None,
233 }
234 }
235 _ => None,
236 }) else {
237 dcx.emit_err(diagnostics::AutoDiffInvalidApplication { span: item.span() });
238 return ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[item]))vec![item];
239 };
240
241 let meta_item_vec: ThinVec<MetaItemInner> = match meta_item.kind {
242 ast::MetaItemKind::List(ref vec) => vec.clone(),
243 _ => {
244 dcx.emit_err(diagnostics::AutoDiffMissingConfig { span: item.span() });
245 return ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[item]))vec![item];
246 }
247 };
248
249 let has_ret = has_ret(&sig.decl.output);
250
251 let mut ts: Vec<TokenTree> = ::alloc::vec::Vec::new()vec![];
254 if meta_item_vec.is_empty() {
255 dcx.emit_err(diagnostics::AutoDiffMissingConfig { span: item.span() });
257 return ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[item]))vec![item];
258 }
259
260 let mode_symbol = match mode {
261 DiffMode::Forward => sym::Forward,
262 DiffMode::Reverse => sym::Reverse,
263 _ => {
::core::panicking::panic_fmt(format_args!("internal error: entered unreachable code: {0}",
format_args!("Unsupported mode: {0:?}", mode)));
}unreachable!("Unsupported mode: {:?}", mode),
264 };
265
266 let mode_token =
268 Token::new(TokenKind::Ident(mode_symbol, IdentKind::Normal), Span::default());
269 ts.insert(0, TokenTree::Token(mode_token, Spacing::Joint));
270 ts.insert(
271 1,
272 TokenTree::Token(Token::new(TokenKind::Comma, Span::default()), Spacing::Alone),
273 );
274
275 let start_position;
278 let kind: LitKind = LitKind::Integer;
279 let symbol;
280 if meta_item_vec.len() >= 2
281 && let Some(width) = width(&meta_item_vec[1])
282 {
283 start_position = 2;
284 symbol = Symbol::intern(&width.to_string());
285 } else {
286 start_position = 1;
287 symbol = sym::integer(1);
288 }
289
290 let l: Lit = Lit { kind, symbol, suffix: None };
291 let t = Token::new(TokenKind::Literal(l), Span::default());
292 let comma = Token::new(TokenKind::Comma, Span::default());
293 ts.push(TokenTree::Token(t, Spacing::Joint));
294 ts.push(TokenTree::Token(comma, Spacing::Alone));
295
296 for t in meta_item_vec.clone()[start_position..].iter() {
297 meta_item_inner_to_ts(t, &mut ts);
298 }
299
300 if !has_ret {
301 let t = Token::new(TokenKind::Ident(sym::None, IdentKind::Normal), Span::default());
304 ts.push(TokenTree::Token(t, Spacing::Joint));
305 ts.push(TokenTree::Token(comma, Spacing::Alone));
306 }
307 ts.pop();
309 let ts: TokenStream = TokenStream::from_iter(ts);
310
311 let x: RustcAutodiff = from_ast(ecx, &meta_item_vec, has_ret, mode);
312 if !x.is_active() {
313 return ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[item]))vec![item];
316 }
317 let span = ecx.with_def_site_ctxt(expand_span);
318
319 let d_sig = gen_enzyme_decl(ecx, &sig, &x, span);
320
321 let d_body = ecx.block(
322 span,
323 {
let len = [()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(call_autodiff(ecx, primal, first_ident(&meta_item_vec[0]), span,
&sig, &d_sig, &generics, is_impl));
vec
}thin_vec![call_autodiff(
324 ecx,
325 primal,
326 first_ident(&meta_item_vec[0]),
327 span,
328 &sig,
329 &d_sig,
330 &generics,
331 is_impl,
332 )],
333 );
334
335 let d_fn = Box::new(ast::Fn {
337 defaultness: ast::Defaultness::Implicit,
338 sig: d_sig,
339 ident: first_ident(&meta_item_vec[0]),
340 generics,
341 contract: None,
342 body: Some(d_body),
343 define_opaque: None,
344 eii_impl: None,
345 });
346 let mut rustc_ad_attr =
347 Box::new(ast::NormalAttr::from_ident(Ident::with_dummy_span(sym::rustc_autodiff)));
348
349 let ts2: Vec<TokenTree> = ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[TokenTree::Token(Token::new(TokenKind::Ident(sym::never,
IdentKind::Normal), span), Spacing::Joint)]))vec![TokenTree::Token(
350 Token::new(TokenKind::Ident(sym::never, IdentKind::Normal), span),
351 Spacing::Joint,
352 )];
353 let never_arg = ast::DelimArgs {
354 dspan: DelimSpan::from_single(span),
355 delim: ast::token::Delimiter::Parenthesis,
356 tokens: TokenStream::from_iter(ts2),
357 };
358 let inline_item = ast::AttrItem {
359 unsafety: ast::Safety::Default,
360 path: ast::Path::from_ident(Ident::with_dummy_span(sym::inline)),
361 args: ast::AttrArgs::Delimited(never_arg),
362 span: DUMMY_SP,
363 };
364 let inline_never_attr = Box::new(ast::NormalAttr { item: inline_item, tokens: None });
365 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
366 let attr = outer_normal_attr(&rustc_ad_attr, new_id, span);
367 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
368 let inline_never = outer_normal_attr(&inline_never_attr, new_id, span);
369
370 fn same_attribute(attr: &ast::AttrKind, item: &ast::AttrKind) -> bool {
372 match (attr, item) {
373 (ast::AttrKind::Normal(a), ast::AttrKind::Normal(b)) => {
374 let a = &a.item.path;
375 let b = &b.item.path;
376 a.segments.iter().eq_by(&b.segments, |a, b| a.ident == b.ident)
377 }
378 _ => false,
379 }
380 }
381
382 let mut has_inline_never = false;
383
384 let orig_annotatable: Annotatable = match item {
386 Annotatable::Item(ref mut iitem) => {
387 if !iitem.attrs.iter().any(|a| same_attribute(&a.kind, &attr.kind)) {
388 iitem.attrs.push(attr);
389 }
390 if iitem.attrs.iter().any(|a| same_attribute(&a.kind, &inline_never.kind)) {
391 has_inline_never = true;
392 }
393 Annotatable::Item(iitem.clone())
394 }
395 Annotatable::AssocItem(ref mut assoc_item, ctxt @ (Impl { .. } | Trait)) => {
396 if !assoc_item.attrs.iter().any(|a| same_attribute(&a.kind, &attr.kind)) {
397 assoc_item.attrs.push(attr);
398 }
399 if assoc_item.attrs.iter().any(|a| same_attribute(&a.kind, &inline_never.kind)) {
400 has_inline_never = true;
401 }
402 Annotatable::AssocItem(assoc_item.clone(), ctxt)
403 }
404 Annotatable::Stmt(ref mut stmt) => {
405 match stmt.kind {
406 ast::StmtKind::Item(ref mut iitem) => {
407 if !iitem.attrs.iter().any(|a| same_attribute(&a.kind, &attr.kind)) {
408 iitem.attrs.push(attr);
409 }
410 if iitem.attrs.iter().any(|a| same_attribute(&a.kind, &inline_never.kind)) {
411 has_inline_never = true;
412 }
413 }
414 _ => {
::core::panicking::panic_fmt(format_args!("internal error: entered unreachable code: {0}",
format_args!("stmt kind checked previously")));
}unreachable!("stmt kind checked previously"),
415 };
416
417 Annotatable::Stmt(stmt.clone())
418 }
419 _ => {
420 {
::core::panicking::panic_fmt(format_args!("internal error: entered unreachable code: {0}",
format_args!("annotatable kind checked previously")));
}unreachable!("annotatable kind checked previously")
421 }
422 };
423 rustc_ad_attr.item.args = rustc_ast::AttrArgs::Delimited(rustc_ast::DelimArgs {
425 dspan: DelimSpan::dummy(),
426 delim: rustc_ast::token::Delimiter::Parenthesis,
427 tokens: ts,
428 });
429
430 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
431 let d_attr = outer_normal_attr(&rustc_ad_attr, new_id, span);
432
433 let mut d_attrs = {
let len = [()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(d_attr);
vec
}thin_vec![d_attr];
435
436 if has_inline_never {
437 d_attrs.push(inline_never);
438 }
439
440 let d_annotatable = match &item {
441 Annotatable::AssocItem(_, ctxt) => {
442 let assoc_item: AssocItemKind = ast::AssocItemKind::Fn(d_fn);
443 let d_fn = Box::new(ast::AssocItem {
444 attrs: d_attrs,
445 id: ast::DUMMY_NODE_ID,
446 span,
447 vis,
448 kind: assoc_item,
449 tokens: None,
450 });
451 Annotatable::AssocItem(d_fn, *ctxt)
452 }
453 Annotatable::Item(_) => {
454 let mut d_fn = ecx.item(span, d_attrs, ItemKind::Fn(d_fn));
455 d_fn.vis = vis;
456
457 Annotatable::Item(d_fn)
458 }
459 Annotatable::Stmt(_) => {
460 let mut d_fn = ecx.item(span, d_attrs, ItemKind::Fn(d_fn));
461 d_fn.vis = vis;
462
463 Annotatable::Stmt(Box::new(ast::Stmt {
464 id: ast::DUMMY_NODE_ID,
465 kind: ast::StmtKind::Item(d_fn),
466 span,
467 }))
468 }
469 _ => {
470 {
::core::panicking::panic_fmt(format_args!("internal error: entered unreachable code: {0}",
format_args!("item kind checked previously")));
}unreachable!("item kind checked previously")
471 }
472 };
473
474 ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[orig_annotatable, d_annotatable]))vec![orig_annotatable, d_annotatable]
475 }
476
477 fn assure_mut_ref(ty: &ast::Ty) -> ast::Ty {
480 let mut ty = ty.clone();
481 match ty.kind {
482 TyKind::Ptr(ref mut mut_ty) => {
483 mut_ty.mutbl = ast::Mutability::Mut;
484 }
485 TyKind::Ref(_, ref mut mut_ty) => {
486 mut_ty.mutbl = ast::Mutability::Mut;
487 }
488 _ => {
489 {
::core::panicking::panic_fmt(format_args!("unsupported type: {0:?}", ty));
};panic!("unsupported type: {:?}", ty);
490 }
491 }
492 ty
493 }
494
495 fn call_autodiff(
500 ecx: &ExtCtxt<'_>,
501 primal: Ident,
502 diff: Ident,
503 span: Span,
504 p_sig: &FnSig,
505 d_sig: &FnSig,
506 generics: &Generics,
507 is_impl: bool,
508 ) -> rustc_ast::Stmt {
509 let primal_path_expr = gen_turbofish_expr(ecx, primal, generics, span, is_impl);
510
511 let self_ty = || ecx.ty_path(ast::Path::from_ident(Ident::with_dummy_span(kw::SelfUpper)));
512 let fn_ptr_params: ThinVec<ast::Param> = p_sig
513 .decl
514 .inputs
515 .iter()
516 .map(|param| {
517 let ty = match ¶m.ty.kind {
518 TyKind::ImplicitSelf => self_ty(),
519 TyKind::Ref(lt, mt) if #[allow(non_exhaustive_omitted_patterns)] match mt.ty.kind {
TyKind::ImplicitSelf => true,
_ => false,
}matches!(mt.ty.kind, TyKind::ImplicitSelf) => ecx
520 .ty(span, TyKind::Ref(*lt, ast::MutTy { ty: self_ty(), mutbl: mt.mutbl })),
521 TyKind::Ptr(mt) if #[allow(non_exhaustive_omitted_patterns)] match mt.ty.kind {
TyKind::ImplicitSelf => true,
_ => false,
}matches!(mt.ty.kind, TyKind::ImplicitSelf) => {
522 ecx.ty(span, TyKind::Ptr(ast::MutTy { ty: self_ty(), mutbl: mt.mutbl }))
523 }
524 _ => param.ty.clone(),
525 };
526 ast::Param {
527 attrs: ast::AttrVec::new(),
528 ty,
529 pat: Box::new(ecx.pat_wild(span)),
530 id: ast::DUMMY_NODE_ID,
531 span,
532 is_placeholder: false,
533 }
534 })
535 .collect();
536 let fn_ptr_ty = ecx.ty(
537 span,
538 TyKind::FnPtr(Box::new(ast::FnPtrTy {
539 safety: p_sig.header.safety,
540 ext: p_sig.header.ext,
541 generic_params: ThinVec::new(),
542 decl: Box::new(ast::FnDecl {
543 inputs: fn_ptr_params,
544 output: p_sig.decl.output.clone(),
545 }),
546 decl_span: span,
547 })),
548 );
549 let primal_fn_ptr = ecx.expr(span, ast::ExprKind::Cast(primal_path_expr, fn_ptr_ty));
550
551 let diff_path_expr = gen_turbofish_expr(ecx, diff, generics, span, is_impl);
552
553 let tuple_expr = ecx.expr_tuple(
554 span,
555 d_sig
556 .decl
557 .inputs
558 .iter()
559 .map(|arg| match arg.pat.kind {
560 PatKind::Ident(_, ident, _) => ecx.expr_path(ecx.path_ident(span, ident)),
561 _ => ::core::panicking::panic("not implemented")unimplemented!(),
562 })
563 .collect::<ThinVec<_>>(),
564 );
565
566 let enzyme_path_idents = ecx.std_path(&[sym::intrinsics, sym::autodiff]);
567 let enzyme_path = ecx.path(span, enzyme_path_idents);
568 let call_expr = ecx.expr_call(
569 span,
570 ecx.expr_path(enzyme_path),
571 ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[primal_fn_ptr, diff_path_expr, tuple_expr]))vec![primal_fn_ptr, diff_path_expr, tuple_expr].into(),
572 );
573
574 ecx.stmt_expr(call_expr)
575 }
576
577 fn gen_turbofish_expr(
581 ecx: &ExtCtxt<'_>,
582 ident: Ident,
583 generics: &Generics,
584 span: Span,
585 is_impl: bool,
586 ) -> Box<ast::Expr> {
587 let generic_args = generics
588 .params
589 .iter()
590 .filter_map(|p| match &p.kind {
591 GenericParamKind::Type { .. } => {
592 let path = ast::Path::from_ident(p.ident);
593 let ty = ecx.ty_path(path);
594 Some(AngleBracketedArg::Arg(GenericArg::Type(ty)))
595 }
596 GenericParamKind::Const { .. } => {
597 let expr = ecx.expr_path(ast::Path::from_ident(p.ident));
598 let anon_const = AnonConst { id: ast::DUMMY_NODE_ID, value: expr };
599 Some(AngleBracketedArg::Arg(GenericArg::Const(anon_const)))
600 }
601 GenericParamKind::Lifetime => None,
602 })
603 .collect::<ThinVec<_>>();
604
605 let args: AngleBracketedArgs = AngleBracketedArgs { span, args: generic_args };
606
607 let segment = PathSegment {
608 ident,
609 id: ast::DUMMY_NODE_ID,
610 args: Some(Box::new(GenericArgs::AngleBracketed(args))),
611 };
612
613 let segments = if is_impl {
614 {
let len = [(), ()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(PathSegment {
ident: Ident::from_str("Self"),
id: ast::DUMMY_NODE_ID,
args: None,
});
vec.push(segment);
vec
}thin_vec![
615 PathSegment { ident: Ident::from_str("Self"), id: ast::DUMMY_NODE_ID, args: None },
616 segment,
617 ]
618 } else {
619 {
let len = [()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(segment);
vec
}thin_vec![segment]
620 };
621
622 let path = Path { span, segments };
623
624 ecx.expr_path(path)
625 }
626
627 fn gen_enzyme_decl(
639 ecx: &ExtCtxt<'_>,
640 sig: &ast::FnSig,
641 x: &RustcAutodiff,
642 span: Span,
643 ) -> ast::FnSig {
644 let dcx = ecx.sess.dcx();
645 let has_ret = has_ret(&sig.decl.output);
646 let sig_args = sig.decl.inputs.len() + if has_ret { 1 } else { 0 };
647 let num_activities = x.input_activity.len() + if x.has_ret_activity() { 1 } else { 0 };
648 if sig_args != num_activities {
649 dcx.emit_err(diagnostics::AutoDiffInvalidNumberActivities {
650 span,
651 expected: sig_args,
652 found: num_activities,
653 });
654 return sig.clone();
656 }
657 if !(sig.decl.inputs.len() == x.input_activity.len()) {
::core::panicking::panic("assertion failed: sig.decl.inputs.len() == x.input_activity.len()")
};assert!(sig.decl.inputs.len() == x.input_activity.len());
658 if !(has_ret == x.has_ret_activity()) {
::core::panicking::panic("assertion failed: has_ret == x.has_ret_activity()")
};assert!(has_ret == x.has_ret_activity());
659 let mut d_decl = sig.decl.clone();
660 let mut d_inputs = Vec::new();
661 let mut new_inputs = Vec::new();
662 let mut idents = Vec::new();
663 let mut act_ret = ThinVec::new();
664
665 let mut errors = false;
668 for (arg, activity) in sig.decl.inputs.iter().zip(x.input_activity.iter()) {
669 if !valid_input_activity(x.mode, *activity) {
670 dcx.emit_err(diagnostics::AutoDiffInvalidApplicationModeAct {
671 span,
672 mode: x.mode.to_string(),
673 act: activity.to_string(),
674 });
675 errors = true;
676 }
677 if !valid_ty_for_activity(&arg.ty, *activity) {
678 dcx.emit_err(diagnostics::AutoDiffInvalidTypeForActivity {
679 span: arg.ty.span,
680 act: activity.to_string(),
681 });
682 errors = true;
683 }
684 }
685
686 if has_ret && !valid_ret_activity(x.mode, x.ret_activity) {
687 dcx.emit_err(diagnostics::AutoDiffInvalidRetAct {
688 span,
689 mode: x.mode.to_string(),
690 act: x.ret_activity.to_string(),
691 });
692 }
695
696 if errors {
697 return sig.clone();
699 }
700
701 let unsafe_activities = x
702 .input_activity
703 .iter()
704 .any(|&act| #[allow(non_exhaustive_omitted_patterns)] match act {
DiffActivity::DuplicatedOnly | DiffActivity::DualOnly => true,
_ => false,
}matches!(act, DiffActivity::DuplicatedOnly | DiffActivity::DualOnly));
705 for (arg, activity) in sig.decl.inputs.iter().zip(x.input_activity.iter()) {
706 d_inputs.push(arg.clone());
707 match activity {
708 DiffActivity::Active => {
709 act_ret.push(arg.ty.clone());
710 }
712 DiffActivity::ActiveOnly => {
713 }
716 DiffActivity::Duplicated | DiffActivity::DuplicatedOnly => {
717 for i in 0..x.width {
718 let mut shadow_arg = arg.clone();
719 *shadow_arg.ty = assure_mut_ref(&arg.ty);
721 let old_name = if let PatKind::Ident(_, ident, _) = arg.pat.kind {
722 ident.name
723 } else {
724 {
use ::tracing::__macro_support::Callsite as _;
static __CALLSITE: ::tracing::callsite::DefaultCallsite =
{
static META: ::tracing::Metadata<'static> =
{
::tracing_core::metadata::Metadata::new("event /rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs:724",
"rustc_builtin_macros::autodiff::llvm_enzyme",
::tracing::Level::DEBUG,
::tracing_core::__macro_support::Option::Some("/rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs"),
::tracing_core::__macro_support::Option::Some(724u32),
::tracing_core::__macro_support::Option::Some("rustc_builtin_macros::autodiff::llvm_enzyme"),
::tracing_core::field::FieldSet::new(&["message"],
::tracing_core::callsite::Identifier(&__CALLSITE)),
::tracing::metadata::Kind::EVENT)
};
::tracing::callsite::DefaultCallsite::new(&META)
};
let enabled =
::tracing::Level::DEBUG <= ::tracing::level_filters::STATIC_MAX_LEVEL
&&
::tracing::Level::DEBUG <=
::tracing::level_filters::LevelFilter::current() &&
{
let interest = __CALLSITE.interest();
!interest.is_never() &&
::tracing::__macro_support::__is_enabled(__CALLSITE.metadata(),
interest)
};
if enabled {
(|value_set: ::tracing::field::ValueSet|
{
let meta = __CALLSITE.metadata();
::tracing::Event::dispatch(meta, &value_set);
;
})({
#[allow(unused_imports)]
use ::tracing::field::{debug, display, Value};
__CALLSITE.metadata().fields().value_set_all(&[(::tracing::__macro_support::Option::Some(&format_args!("{0:#?}",
&shadow_arg.pat) as &dyn ::tracing::field::Value))])
});
} else { ; }
};debug!("{:#?}", &shadow_arg.pat);
725 { ::core::panicking::panic_fmt(format_args!("not an ident?")); };panic!("not an ident?");
726 };
727 let name: String = ::alloc::__export::must_use({
::alloc::fmt::format(format_args!("d{0}_{1}", old_name, i))
})format!("d{}_{}", old_name, i);
728 new_inputs.push(name.clone());
729 let ident = Ident::from_str_and_span(&name, shadow_arg.pat.span);
730 *shadow_arg.pat = ast::Pat {
731 id: ast::DUMMY_NODE_ID,
732 kind: PatKind::Ident(BindingMode::NONE, ident, None),
733 span: shadow_arg.pat.span,
734 };
735 d_inputs.push(shadow_arg.clone());
736 }
737 }
738 DiffActivity::Dual
739 | DiffActivity::DualOnly
740 | DiffActivity::Dualv
741 | DiffActivity::DualvOnly => {
742 let iterations =
745 if #[allow(non_exhaustive_omitted_patterns)] match activity {
DiffActivity::Dualv | DiffActivity::DualvOnly => true,
_ => false,
}matches!(activity, DiffActivity::Dualv | DiffActivity::DualvOnly) {
746 1
747 } else {
748 x.width
749 };
750 for i in 0..iterations {
751 let mut shadow_arg = arg.clone();
752 let old_name = if let PatKind::Ident(_, ident, _) = arg.pat.kind {
753 ident.name
754 } else {
755 {
use ::tracing::__macro_support::Callsite as _;
static __CALLSITE: ::tracing::callsite::DefaultCallsite =
{
static META: ::tracing::Metadata<'static> =
{
::tracing_core::metadata::Metadata::new("event /rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs:755",
"rustc_builtin_macros::autodiff::llvm_enzyme",
::tracing::Level::DEBUG,
::tracing_core::__macro_support::Option::Some("/rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs"),
::tracing_core::__macro_support::Option::Some(755u32),
::tracing_core::__macro_support::Option::Some("rustc_builtin_macros::autodiff::llvm_enzyme"),
::tracing_core::field::FieldSet::new(&["message"],
::tracing_core::callsite::Identifier(&__CALLSITE)),
::tracing::metadata::Kind::EVENT)
};
::tracing::callsite::DefaultCallsite::new(&META)
};
let enabled =
::tracing::Level::DEBUG <= ::tracing::level_filters::STATIC_MAX_LEVEL
&&
::tracing::Level::DEBUG <=
::tracing::level_filters::LevelFilter::current() &&
{
let interest = __CALLSITE.interest();
!interest.is_never() &&
::tracing::__macro_support::__is_enabled(__CALLSITE.metadata(),
interest)
};
if enabled {
(|value_set: ::tracing::field::ValueSet|
{
let meta = __CALLSITE.metadata();
::tracing::Event::dispatch(meta, &value_set);
;
})({
#[allow(unused_imports)]
use ::tracing::field::{debug, display, Value};
__CALLSITE.metadata().fields().value_set_all(&[(::tracing::__macro_support::Option::Some(&format_args!("{0:#?}",
&shadow_arg.pat) as &dyn ::tracing::field::Value))])
});
} else { ; }
};debug!("{:#?}", &shadow_arg.pat);
756 { ::core::panicking::panic_fmt(format_args!("not an ident?")); };panic!("not an ident?");
757 };
758 let name: String = ::alloc::__export::must_use({
::alloc::fmt::format(format_args!("b{0}_{1}", old_name, i))
})format!("b{}_{}", old_name, i);
759 new_inputs.push(name.clone());
760 let ident = Ident::from_str_and_span(&name, shadow_arg.pat.span);
761 *shadow_arg.pat = ast::Pat {
762 id: ast::DUMMY_NODE_ID,
763 kind: PatKind::Ident(BindingMode::NONE, ident, None),
764 span: shadow_arg.pat.span,
765 };
766 d_inputs.push(shadow_arg.clone());
767 }
768 }
769 DiffActivity::Const => {
770 }
772 DiffActivity::None | DiffActivity::FakeActivitySize(_) => {
773 { ::core::panicking::panic_fmt(format_args!("Should not happen")); };panic!("Should not happen");
774 }
775 }
776 if let PatKind::Ident(_, ident, _) = arg.pat.kind {
777 idents.push(ident);
778 } else {
779 { ::core::panicking::panic_fmt(format_args!("not an ident?")); };panic!("not an ident?");
780 }
781 }
782
783 let active_only_ret = x.ret_activity == DiffActivity::ActiveOnly;
784 if active_only_ret {
785 if !x.mode.is_rev() {
::core::panicking::panic("assertion failed: x.mode.is_rev()")
};assert!(x.mode.is_rev());
786 }
787
788 if x.mode.is_rev() {
791 match x.ret_activity {
792 DiffActivity::Active | DiffActivity::ActiveOnly => {
793 let ty = match d_decl.output {
794 FnRetTy::Ty(ref ty) => ty.clone(),
795 FnRetTy::Default(span) => {
796 {
::core::panicking::panic_fmt(format_args!("Did not expect Default ret ty: {0:?}",
span));
};panic!("Did not expect Default ret ty: {:?}", span);
797 }
798 };
799 let name = "dret".to_string();
800 let ident = Ident::from_str_and_span(&name, ty.span);
801 let shadow_arg = ast::Param {
802 attrs: ThinVec::new(),
803 ty: ty.clone(),
804 pat: Box::new(ast::Pat {
805 id: ast::DUMMY_NODE_ID,
806 kind: PatKind::Ident(BindingMode::NONE, ident, None),
807 span: ty.span,
808 }),
809 id: ast::DUMMY_NODE_ID,
810 span: ty.span,
811 is_placeholder: false,
812 };
813 d_inputs.push(shadow_arg);
814 new_inputs.push(name);
815 }
816 _ => {}
817 }
818 }
819 d_decl.inputs = d_inputs.into();
820
821 if x.mode.is_fwd() {
822 let ty = match d_decl.output {
823 FnRetTy::Ty(ref ty) => ty.clone(),
824 FnRetTy::Default(span) => {
825 let kind = TyKind::Tup(ThinVec::new());
827 let ty = Box::new(rustc_ast::Ty { kind, id: ast::DUMMY_NODE_ID, span });
828 d_decl.output = FnRetTy::Ty(ty.clone());
829 if !#[allow(non_exhaustive_omitted_patterns)] match x.ret_activity {
DiffActivity::None => true,
_ => false,
} {
::core::panicking::panic("assertion failed: matches!(x.ret_activity, DiffActivity::None)")
};assert!(matches!(x.ret_activity, DiffActivity::None));
830 ty
832 }
833 };
834
835 if #[allow(non_exhaustive_omitted_patterns)] match x.ret_activity {
DiffActivity::Dual | DiffActivity::Dualv => true,
_ => false,
}matches!(x.ret_activity, DiffActivity::Dual | DiffActivity::Dualv) {
836 let kind = if x.width == 1 || #[allow(non_exhaustive_omitted_patterns)] match x.ret_activity {
DiffActivity::Dualv => true,
_ => false,
}matches!(x.ret_activity, DiffActivity::Dualv) {
837 TyKind::Tup({
let len = [(), ()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(ty.clone());
vec.push(ty.clone());
vec
}thin_vec![ty.clone(), ty.clone()])
840 } else {
841 let anon_const = rustc_ast::AnonConst {
843 id: ast::DUMMY_NODE_ID,
844 value: ecx.expr_usize(span, 1 + x.width as usize),
845 };
846 TyKind::Array(ty.clone(), anon_const)
847 };
848 let ty = Box::new(rustc_ast::Ty { kind, id: ty.id, span: ty.span });
849 d_decl.output = FnRetTy::Ty(ty);
850 }
851 if #[allow(non_exhaustive_omitted_patterns)] match x.ret_activity {
DiffActivity::DualOnly | DiffActivity::DualvOnly => true,
_ => false,
}matches!(x.ret_activity, DiffActivity::DualOnly | DiffActivity::DualvOnly) {
852 if x.width > 1 {
856 let anon_const = rustc_ast::AnonConst {
857 id: ast::DUMMY_NODE_ID,
858 value: ecx.expr_usize(span, x.width as usize),
859 };
860 let kind = TyKind::Array(ty.clone(), anon_const);
861 let ty = Box::new(rustc_ast::Ty { kind, id: ty.id, span: ty.span });
862 d_decl.output = FnRetTy::Ty(ty);
863 }
864 }
865 }
866
867 d_decl.output =
869 if active_only_ret { FnRetTy::Default(span) } else { d_decl.output.clone() };
870
871 {
use ::tracing::__macro_support::Callsite as _;
static __CALLSITE: ::tracing::callsite::DefaultCallsite =
{
static META: ::tracing::Metadata<'static> =
{
::tracing_core::metadata::Metadata::new("event /rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs:871",
"rustc_builtin_macros::autodiff::llvm_enzyme",
::tracing::Level::TRACE,
::tracing_core::__macro_support::Option::Some("/rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs"),
::tracing_core::__macro_support::Option::Some(871u32),
::tracing_core::__macro_support::Option::Some("rustc_builtin_macros::autodiff::llvm_enzyme"),
::tracing_core::field::FieldSet::new(&["message"],
::tracing_core::callsite::Identifier(&__CALLSITE)),
::tracing::metadata::Kind::EVENT)
};
::tracing::callsite::DefaultCallsite::new(&META)
};
let enabled =
::tracing::Level::TRACE <= ::tracing::level_filters::STATIC_MAX_LEVEL
&&
::tracing::Level::TRACE <=
::tracing::level_filters::LevelFilter::current() &&
{
let interest = __CALLSITE.interest();
!interest.is_never() &&
::tracing::__macro_support::__is_enabled(__CALLSITE.metadata(),
interest)
};
if enabled {
(|value_set: ::tracing::field::ValueSet|
{
let meta = __CALLSITE.metadata();
::tracing::Event::dispatch(meta, &value_set);
;
})({
#[allow(unused_imports)]
use ::tracing::field::{debug, display, Value};
__CALLSITE.metadata().fields().value_set_all(&[(::tracing::__macro_support::Option::Some(&format_args!("act_ret: {0:?}",
act_ret) as &dyn ::tracing::field::Value))])
});
} else { ; }
};trace!("act_ret: {:?}", act_ret);
872
873 if act_ret.len() > 0 {
877 let ret_ty = match d_decl.output {
878 FnRetTy::Ty(ref ty) => {
879 if !active_only_ret {
880 act_ret.insert(0, ty.clone());
881 }
882 let kind = TyKind::Tup(act_ret);
883 Box::new(rustc_ast::Ty { kind, id: ty.id, span: ty.span })
884 }
885 FnRetTy::Default(span) => {
886 if act_ret.len() == 1 {
887 act_ret[0].clone()
888 } else {
889 let kind = TyKind::Tup(act_ret);
890 Box::new(rustc_ast::Ty { kind, id: ast::DUMMY_NODE_ID, span })
891 }
892 }
893 };
894 d_decl.output = FnRetTy::Ty(ret_ty);
895 }
896
897 let mut d_header = sig.header;
898 if unsafe_activities {
899 d_header.safety = rustc_ast::Safety::Unsafe(span);
900 }
901 let d_sig = FnSig { header: d_header, decl: d_decl, span };
902 {
use ::tracing::__macro_support::Callsite as _;
static __CALLSITE: ::tracing::callsite::DefaultCallsite =
{
static META: ::tracing::Metadata<'static> =
{
::tracing_core::metadata::Metadata::new("event /rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs:902",
"rustc_builtin_macros::autodiff::llvm_enzyme",
::tracing::Level::TRACE,
::tracing_core::__macro_support::Option::Some("/rustc-dev/6bb1652a020e80cef79332741d89e996d71933c9/compiler/rustc_builtin_macros/src/autodiff.rs"),
::tracing_core::__macro_support::Option::Some(902u32),
::tracing_core::__macro_support::Option::Some("rustc_builtin_macros::autodiff::llvm_enzyme"),
::tracing_core::field::FieldSet::new(&["message"],
::tracing_core::callsite::Identifier(&__CALLSITE)),
::tracing::metadata::Kind::EVENT)
};
::tracing::callsite::DefaultCallsite::new(&META)
};
let enabled =
::tracing::Level::TRACE <= ::tracing::level_filters::STATIC_MAX_LEVEL
&&
::tracing::Level::TRACE <=
::tracing::level_filters::LevelFilter::current() &&
{
let interest = __CALLSITE.interest();
!interest.is_never() &&
::tracing::__macro_support::__is_enabled(__CALLSITE.metadata(),
interest)
};
if enabled {
(|value_set: ::tracing::field::ValueSet|
{
let meta = __CALLSITE.metadata();
::tracing::Event::dispatch(meta, &value_set);
;
})({
#[allow(unused_imports)]
use ::tracing::field::{debug, display, Value};
__CALLSITE.metadata().fields().value_set_all(&[(::tracing::__macro_support::Option::Some(&format_args!("Generated signature: {0:?}",
d_sig) as &dyn ::tracing::field::Value))])
});
} else { ; }
};trace!("Generated signature: {:?}", d_sig);
903 d_sig
904 }
905}
906
907pub(crate) use llvm_enzyme::{expand_forward, expand_reverse};