1mod llvm_enzyme {
7 use std::str::FromStr;
8 use std::string::String;
9
10 use rustc_ast::expand::autodiff_attrs::{
11 AutoDiffAttrs, DiffActivity, DiffMode, valid_input_activity, valid_ret_activity,
12 valid_ty_for_activity,
13 };
14 use rustc_ast::ptr::P;
15 use rustc_ast::token::{Token, TokenKind};
16 use rustc_ast::tokenstream::*;
17 use rustc_ast::visit::AssocCtxt::*;
18 use rustc_ast::{
19 self as ast, AssocItemKind, BindingMode, FnRetTy, FnSig, Generics, ItemKind, MetaItemInner,
20 PatKind, TyKind,
21 };
22 use rustc_expand::base::{Annotatable, ExtCtxt};
23 use rustc_span::{Ident, Span, Symbol, kw, sym};
24 use thin_vec::{ThinVec, thin_vec};
25 use tracing::{debug, trace};
26
27 use crate::errors;
28
29 pub(crate) fn outer_normal_attr(
30 kind: &P<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 let segments = &x.meta_item().unwrap().path.segments;
49 assert!(segments.len() == 1);
50 segments[0].ident
51 }
52
53 fn name(x: &MetaItemInner) -> String {
54 first_ident(x).name.to_string()
55 }
56
57 pub(crate) fn from_ast(
58 ecx: &mut ExtCtxt<'_>,
59 meta_item: &ThinVec<MetaItemInner>,
60 has_ret: bool,
61 ) -> AutoDiffAttrs {
62 let dcx = ecx.sess.dcx();
63 let mode = name(&meta_item[1]);
64 let Ok(mode) = DiffMode::from_str(&mode) else {
65 dcx.emit_err(errors::AutoDiffInvalidMode { span: meta_item[1].span(), mode });
66 return AutoDiffAttrs::error();
67 };
68 let mut activities: Vec<DiffActivity> = vec![];
69 let mut errors = false;
70 for x in &meta_item[2..] {
71 let activity_str = name(&x);
72 let res = DiffActivity::from_str(&activity_str);
73 match res {
74 Ok(x) => activities.push(x),
75 Err(_) => {
76 dcx.emit_err(errors::AutoDiffUnknownActivity {
77 span: x.span(),
78 act: activity_str,
79 });
80 errors = true;
81 }
82 };
83 }
84 if errors {
85 return AutoDiffAttrs::error();
86 }
87
88 let (ret_activity, input_activity) = if has_ret {
91 let Some((last, rest)) = activities.split_last() else {
92 unreachable!(
93 "should not be reachable because we counted the number of activities previously"
94 );
95 };
96 (last, rest)
97 } else {
98 (&DiffActivity::None, activities.as_slice())
99 };
100
101 AutoDiffAttrs { mode, ret_activity: *ret_activity, input_activity: input_activity.to_vec() }
102 }
103
104 pub(crate) fn expand(
138 ecx: &mut ExtCtxt<'_>,
139 expand_span: Span,
140 meta_item: &ast::MetaItem,
141 mut item: Annotatable,
142 ) -> Vec<Annotatable> {
143 if cfg!(not(llvm_enzyme)) {
144 ecx.sess.dcx().emit_err(errors::AutoDiffSupportNotBuild { span: meta_item.span });
145 return vec![item];
146 }
147 let dcx = ecx.sess.dcx();
148 let (sig, is_impl): (FnSig, bool) = match &item {
150 Annotatable::Item(iitem) => {
151 let sig = match &iitem.kind {
152 ItemKind::Fn(box ast::Fn { sig, .. }) => sig,
153 _ => {
154 dcx.emit_err(errors::AutoDiffInvalidApplication { span: item.span() });
155 return vec![item];
156 }
157 };
158 (sig.clone(), false)
159 }
160 Annotatable::AssocItem(assoc_item, _) => {
161 let sig = match &assoc_item.kind {
162 ast::AssocItemKind::Fn(box ast::Fn { sig, .. }) => sig,
163 _ => {
164 dcx.emit_err(errors::AutoDiffInvalidApplication { span: item.span() });
165 return vec![item];
166 }
167 };
168 (sig.clone(), true)
169 }
170 _ => {
171 dcx.emit_err(errors::AutoDiffInvalidApplication { span: item.span() });
172 return vec![item];
173 }
174 };
175
176 let meta_item_vec: ThinVec<MetaItemInner> = match meta_item.kind {
177 ast::MetaItemKind::List(ref vec) => vec.clone(),
178 _ => {
179 dcx.emit_err(errors::AutoDiffInvalidApplication { span: item.span() });
180 return vec![item];
181 }
182 };
183
184 let has_ret = has_ret(&sig.decl.output);
185 let sig_span = ecx.with_call_site_ctxt(sig.span);
186
187 let (vis, primal) = match &item {
188 Annotatable::Item(iitem) => (iitem.vis.clone(), iitem.ident.clone()),
189 Annotatable::AssocItem(assoc_item, _) => {
190 (assoc_item.vis.clone(), assoc_item.ident.clone())
191 }
192 _ => {
193 dcx.emit_err(errors::AutoDiffInvalidApplication { span: item.span() });
194 return vec![item];
195 }
196 };
197
198 let comma: Token = Token::new(TokenKind::Comma, Span::default());
201 let mut ts: Vec<TokenTree> = vec![];
202 if meta_item_vec.len() < 2 {
203 dcx.emit_err(errors::AutoDiffMissingConfig { span: item.span() });
206 return vec![item];
207 } else {
208 for t in meta_item_vec.clone()[1..].iter() {
209 let val = first_ident(t);
210 let t = Token::from_ast_ident(val);
211 ts.push(TokenTree::Token(t, Spacing::Joint));
212 ts.push(TokenTree::Token(comma.clone(), Spacing::Alone));
213 }
214 }
215 if !has_ret {
216 let t = Token::new(TokenKind::Ident(sym::None, false.into()), Span::default());
219 ts.push(TokenTree::Token(t, Spacing::Joint));
220 }
221 let ts: TokenStream = TokenStream::from_iter(ts);
222
223 let x: AutoDiffAttrs = from_ast(ecx, &meta_item_vec, has_ret);
224 if !x.is_active() {
225 return vec![item];
228 }
229 let span = ecx.with_def_site_ctxt(expand_span);
230
231 let n_active: u32 = x
232 .input_activity
233 .iter()
234 .filter(|a| **a == DiffActivity::Active || **a == DiffActivity::ActiveOnly)
235 .count() as u32;
236 let (d_sig, new_args, idents, errored) = gen_enzyme_decl(ecx, &sig, &x, span);
237 let d_body = gen_enzyme_body(
238 ecx, &x, n_active, &sig, &d_sig, primal, &new_args, span, sig_span, idents, errored,
239 );
240 let d_ident = first_ident(&meta_item_vec[0]);
241
242 let asdf = Box::new(ast::Fn {
244 defaultness: ast::Defaultness::Final,
245 sig: d_sig,
246 generics: Generics::default(),
247 contract: None,
248 body: Some(d_body),
249 define_opaque: None,
250 });
251 let mut rustc_ad_attr =
252 P(ast::NormalAttr::from_ident(Ident::with_dummy_span(sym::rustc_autodiff)));
253
254 let ts2: Vec<TokenTree> = vec![TokenTree::Token(
255 Token::new(TokenKind::Ident(sym::never, false.into()), span),
256 Spacing::Joint,
257 )];
258 let never_arg = ast::DelimArgs {
259 dspan: ast::tokenstream::DelimSpan::from_single(span),
260 delim: ast::token::Delimiter::Parenthesis,
261 tokens: ast::tokenstream::TokenStream::from_iter(ts2),
262 };
263 let inline_item = ast::AttrItem {
264 unsafety: ast::Safety::Default,
265 path: ast::Path::from_ident(Ident::with_dummy_span(sym::inline)),
266 args: ast::AttrArgs::Delimited(never_arg),
267 tokens: None,
268 };
269 let inline_never_attr = P(ast::NormalAttr { item: inline_item, tokens: None });
270 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
271 let attr = outer_normal_attr(&rustc_ad_attr, new_id, span);
272 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
273 let inline_never = outer_normal_attr(&inline_never_attr, new_id, span);
274
275 fn same_attribute(attr: &ast::AttrKind, item: &ast::AttrKind) -> bool {
277 match (attr, item) {
278 (ast::AttrKind::Normal(a), ast::AttrKind::Normal(b)) => {
279 let a = &a.item.path;
280 let b = &b.item.path;
281 a.segments.len() == b.segments.len()
282 && a.segments.iter().zip(b.segments.iter()).all(|(a, b)| a.ident == b.ident)
283 }
284 _ => false,
285 }
286 }
287
288 let orig_annotatable: Annotatable = match item {
290 Annotatable::Item(ref mut iitem) => {
291 if !iitem.attrs.iter().any(|a| same_attribute(&a.kind, &attr.kind)) {
292 iitem.attrs.push(attr);
293 }
294 if !iitem.attrs.iter().any(|a| same_attribute(&a.kind, &inline_never.kind)) {
295 iitem.attrs.push(inline_never.clone());
296 }
297 Annotatable::Item(iitem.clone())
298 }
299 Annotatable::AssocItem(ref mut assoc_item, i @ Impl) => {
300 if !assoc_item.attrs.iter().any(|a| same_attribute(&a.kind, &attr.kind)) {
301 assoc_item.attrs.push(attr);
302 }
303 if !assoc_item.attrs.iter().any(|a| same_attribute(&a.kind, &inline_never.kind)) {
304 assoc_item.attrs.push(inline_never.clone());
305 }
306 Annotatable::AssocItem(assoc_item.clone(), i)
307 }
308 _ => {
309 unreachable!("annotatable kind checked previously")
310 }
311 };
312 rustc_ad_attr.item.args = rustc_ast::AttrArgs::Delimited(rustc_ast::DelimArgs {
314 dspan: DelimSpan::dummy(),
315 delim: rustc_ast::token::Delimiter::Parenthesis,
316 tokens: ts,
317 });
318 let d_attr = outer_normal_attr(&rustc_ad_attr, new_id, span);
319 let d_annotatable = if is_impl {
320 let assoc_item: AssocItemKind = ast::AssocItemKind::Fn(asdf);
321 let d_fn = P(ast::AssocItem {
322 attrs: thin_vec![d_attr, inline_never],
323 id: ast::DUMMY_NODE_ID,
324 span,
325 vis,
326 ident: d_ident,
327 kind: assoc_item,
328 tokens: None,
329 });
330 Annotatable::AssocItem(d_fn, Impl)
331 } else {
332 let mut d_fn =
333 ecx.item(span, d_ident, thin_vec![d_attr, inline_never], ItemKind::Fn(asdf));
334 d_fn.vis = vis;
335 Annotatable::Item(d_fn)
336 };
337
338 return vec![orig_annotatable, d_annotatable];
339 }
340
341 fn assure_mut_ref(ty: &ast::Ty) -> ast::Ty {
344 let mut ty = ty.clone();
345 match ty.kind {
346 TyKind::Ptr(ref mut mut_ty) => {
347 mut_ty.mutbl = ast::Mutability::Mut;
348 }
349 TyKind::Ref(_, ref mut mut_ty) => {
350 mut_ty.mutbl = ast::Mutability::Mut;
351 }
352 _ => {
353 panic!("unsupported type: {:?}", ty);
354 }
355 }
356 ty
357 }
358
359 fn init_body_helper(
371 ecx: &ExtCtxt<'_>,
372 span: Span,
373 primal: Ident,
374 new_names: &[String],
375 sig_span: Span,
376 new_decl_span: Span,
377 idents: &[Ident],
378 errored: bool,
379 ) -> (P<ast::Block>, P<ast::Expr>, P<ast::Expr>, P<ast::Expr>) {
380 let blackbox_path = ecx.std_path(&[sym::hint, sym::black_box]);
381 let noop = ast::InlineAsm {
382 asm_macro: ast::AsmMacro::Asm,
383 template: vec![ast::InlineAsmTemplatePiece::String("NOP".into())],
384 template_strs: Box::new([]),
385 operands: vec![],
386 clobber_abis: vec![],
387 options: ast::InlineAsmOptions::PURE | ast::InlineAsmOptions::NOMEM,
388 line_spans: vec![],
389 };
390 let noop_expr = ecx.expr_asm(span, P(noop));
391 let unsf = ast::BlockCheckMode::Unsafe(ast::UnsafeSource::CompilerGenerated);
392 let unsf_block = ast::Block {
393 stmts: thin_vec![ecx.stmt_semi(noop_expr)],
394 id: ast::DUMMY_NODE_ID,
395 tokens: None,
396 rules: unsf,
397 span,
398 could_be_bare_literal: false,
399 };
400 let unsf_expr = ecx.expr_block(P(unsf_block));
401 let blackbox_call_expr = ecx.expr_path(ecx.path(span, blackbox_path));
402 let primal_call = gen_primal_call(ecx, span, primal, idents);
403 let black_box_primal_call = ecx.expr_call(
404 new_decl_span,
405 blackbox_call_expr.clone(),
406 thin_vec![primal_call.clone()],
407 );
408 let tup_args = new_names
409 .iter()
410 .map(|arg| ecx.expr_path(ecx.path_ident(span, Ident::from_str(arg))))
411 .collect();
412
413 let black_box_remaining_args = ecx.expr_call(
414 sig_span,
415 blackbox_call_expr.clone(),
416 thin_vec![ecx.expr_tuple(sig_span, tup_args)],
417 );
418
419 let mut body = ecx.block(span, ThinVec::new());
420 body.stmts.push(ecx.stmt_semi(unsf_expr));
421
422 if !errored {
424 body.stmts.push(ecx.stmt_semi(black_box_primal_call.clone()));
425 }
426 body.stmts.push(ecx.stmt_semi(black_box_remaining_args));
427
428 (body, primal_call, black_box_primal_call, blackbox_call_expr)
429 }
430
431 fn gen_enzyme_body(
440 ecx: &ExtCtxt<'_>,
441 x: &AutoDiffAttrs,
442 n_active: u32,
443 sig: &ast::FnSig,
444 d_sig: &ast::FnSig,
445 primal: Ident,
446 new_names: &[String],
447 span: Span,
448 sig_span: Span,
449 idents: Vec<Ident>,
450 errored: bool,
451 ) -> P<ast::Block> {
452 let new_decl_span = d_sig.span;
453
454 let (mut body, primal_call, bb_primal_call, bb_call_expr) = init_body_helper(
463 ecx,
464 span,
465 primal,
466 new_names,
467 sig_span,
468 new_decl_span,
469 &idents,
470 errored,
471 );
472
473 if !has_ret(&d_sig.decl.output) {
474 return body;
476 }
477
478 let primal_ret = has_ret(&sig.decl.output) && !x.has_active_only_ret();
481
482 if primal_ret && n_active == 0 && x.mode.is_rev() {
483 body.stmts.push(ecx.stmt_expr(bb_primal_call));
485 return body;
486 }
487
488 if !primal_ret && n_active == 1 {
489 let ty = match d_sig.decl.output {
491 FnRetTy::Ty(ref ty) => ty.clone(),
492 FnRetTy::Default(span) => {
493 panic!("Did not expect Default ret ty: {:?}", span);
494 }
495 };
496 let arg = ty.kind.is_simple_path().unwrap();
497 let sl: Vec<Symbol> = vec![arg, kw::Default];
498 let tmp = ecx.def_site_path(&sl);
499 let default_call_expr = ecx.expr_path(ecx.path(span, tmp));
500 let default_call_expr = ecx.expr_call(new_decl_span, default_call_expr, thin_vec![]);
501 body.stmts.push(ecx.stmt_expr(default_call_expr));
502 return body;
503 }
504
505 let mut exprs = ThinVec::<P<ast::Expr>>::new();
506 if primal_ret {
507 exprs.push(primal_call);
510 }
511
512 let d_ret_ty = match d_sig.decl.output {
515 FnRetTy::Ty(ref ty) => ty.clone(),
516 FnRetTy::Default(span) => {
517 panic!("Did not expect Default ret ty: {:?}", span);
518 }
519 };
520 let mut d_ret_ty = match d_ret_ty.kind.clone() {
521 TyKind::Tup(ref tys) => tys.clone(),
522 TyKind::Path(_, rustc_ast::Path { segments, .. }) => {
523 if let [segment] = &segments[..]
524 && segment.args.is_none()
525 {
526 let id = vec![segments[0].ident];
527 let kind = TyKind::Path(None, ecx.path(span, id));
528 let ty = P(rustc_ast::Ty { kind, id: ast::DUMMY_NODE_ID, span, tokens: None });
529 thin_vec![ty]
530 } else {
531 panic!("Expected tuple or simple path return type");
532 }
533 }
534 _ => {
535 panic!("Did not expect non-tuple ret ty: {:?}", d_ret_ty);
537 }
538 };
539
540 if x.mode.is_fwd() && x.ret_activity == DiffActivity::Dual {
541 assert!(d_ret_ty.len() == 2);
542 let arg = d_ret_ty[0].kind.is_simple_path().unwrap();
544 let arg2 = d_ret_ty[1].kind.is_simple_path().unwrap();
545 assert!(arg == arg2);
546 let sl: Vec<Symbol> = vec![arg, kw::Default];
547 let tmp = ecx.def_site_path(&sl);
548 let default_call_expr = ecx.expr_path(ecx.path(span, tmp));
549 let default_call_expr = ecx.expr_call(new_decl_span, default_call_expr, thin_vec![]);
550 exprs.push(default_call_expr);
551 } else if x.mode.is_rev() {
552 if primal_ret {
553 d_ret_ty = d_ret_ty[1..].to_vec().into();
555 }
556
557 for arg in d_ret_ty.iter() {
558 let arg = arg.kind.is_simple_path().unwrap();
559 let sl: Vec<Symbol> = vec![arg, kw::Default];
560 let tmp = ecx.def_site_path(&sl);
561 let default_call_expr = ecx.expr_path(ecx.path(span, tmp));
562 let default_call_expr =
563 ecx.expr_call(new_decl_span, default_call_expr, thin_vec![]);
564 exprs.push(default_call_expr);
565 }
566 }
567
568 let ret: P<ast::Expr>;
569 match &exprs[..] {
570 [] => {
571 assert!(!has_ret(&d_sig.decl.output));
572 return body;
574 }
575 [arg] => {
576 ret = ecx.expr_call(new_decl_span, bb_call_expr, thin_vec![arg.clone()]);
577 }
578 args => {
579 let ret_tuple: P<ast::Expr> = ecx.expr_tuple(span, args.into());
580 ret = ecx.expr_call(new_decl_span, bb_call_expr, thin_vec![ret_tuple]);
581 }
582 }
583 assert!(has_ret(&d_sig.decl.output));
584 body.stmts.push(ecx.stmt_expr(ret));
585
586 body
587 }
588
589 fn gen_primal_call(
590 ecx: &ExtCtxt<'_>,
591 span: Span,
592 primal: Ident,
593 idents: &[Ident],
594 ) -> P<ast::Expr> {
595 let has_self = idents.len() > 0 && idents[0].name == kw::SelfLower;
596 if has_self {
597 let args: ThinVec<_> =
598 idents[1..].iter().map(|arg| ecx.expr_path(ecx.path_ident(span, *arg))).collect();
599 let self_expr = ecx.expr_self(span);
600 ecx.expr_method_call(span, self_expr, primal, args)
601 } else {
602 let args: ThinVec<_> =
603 idents.iter().map(|arg| ecx.expr_path(ecx.path_ident(span, *arg))).collect();
604 let primal_call_expr = ecx.expr_path(ecx.path_ident(span, primal));
605 ecx.expr_call(span, primal_call_expr, args)
606 }
607 }
608
609 fn gen_enzyme_decl(
621 ecx: &ExtCtxt<'_>,
622 sig: &ast::FnSig,
623 x: &AutoDiffAttrs,
624 span: Span,
625 ) -> (ast::FnSig, Vec<String>, Vec<Ident>, bool) {
626 let dcx = ecx.sess.dcx();
627 let has_ret = has_ret(&sig.decl.output);
628 let sig_args = sig.decl.inputs.len() + if has_ret { 1 } else { 0 };
629 let num_activities = x.input_activity.len() + if x.has_ret_activity() { 1 } else { 0 };
630 if sig_args != num_activities {
631 dcx.emit_err(errors::AutoDiffInvalidNumberActivities {
632 span,
633 expected: sig_args,
634 found: num_activities,
635 });
636 return (sig.clone(), vec![], vec![], true);
638 }
639 assert!(sig.decl.inputs.len() == x.input_activity.len());
640 assert!(has_ret == x.has_ret_activity());
641 let mut d_decl = sig.decl.clone();
642 let mut d_inputs = Vec::new();
643 let mut new_inputs = Vec::new();
644 let mut idents = Vec::new();
645 let mut act_ret = ThinVec::new();
646
647 let mut errors = false;
650 for (arg, activity) in sig.decl.inputs.iter().zip(x.input_activity.iter()) {
651 if !valid_input_activity(x.mode, *activity) {
652 dcx.emit_err(errors::AutoDiffInvalidApplicationModeAct {
653 span,
654 mode: x.mode.to_string(),
655 act: activity.to_string(),
656 });
657 errors = true;
658 }
659 if !valid_ty_for_activity(&arg.ty, *activity) {
660 dcx.emit_err(errors::AutoDiffInvalidTypeForActivity {
661 span: arg.ty.span,
662 act: activity.to_string(),
663 });
664 errors = true;
665 }
666 }
667
668 if has_ret && !valid_ret_activity(x.mode, x.ret_activity) {
669 dcx.emit_err(errors::AutoDiffInvalidRetAct {
670 span,
671 mode: x.mode.to_string(),
672 act: x.ret_activity.to_string(),
673 });
674 }
677
678 if errors {
679 return (sig.clone(), new_inputs, idents, true);
681 }
682
683 let unsafe_activities = x
684 .input_activity
685 .iter()
686 .any(|&act| matches!(act, DiffActivity::DuplicatedOnly | DiffActivity::DualOnly));
687 for (arg, activity) in sig.decl.inputs.iter().zip(x.input_activity.iter()) {
688 d_inputs.push(arg.clone());
689 match activity {
690 DiffActivity::Active => {
691 act_ret.push(arg.ty.clone());
692 }
693 DiffActivity::ActiveOnly => {
694 }
697 DiffActivity::Duplicated | DiffActivity::DuplicatedOnly => {
698 let mut shadow_arg = arg.clone();
699 shadow_arg.ty = P(assure_mut_ref(&arg.ty));
701 let old_name = if let PatKind::Ident(_, ident, _) = arg.pat.kind {
702 ident.name
703 } else {
704 debug!("{:#?}", &shadow_arg.pat);
705 panic!("not an ident?");
706 };
707 let name: String = format!("d{}", old_name);
708 new_inputs.push(name.clone());
709 let ident = Ident::from_str_and_span(&name, shadow_arg.pat.span);
710 shadow_arg.pat = P(ast::Pat {
711 id: ast::DUMMY_NODE_ID,
712 kind: PatKind::Ident(BindingMode::NONE, ident, None),
713 span: shadow_arg.pat.span,
714 tokens: shadow_arg.pat.tokens.clone(),
715 });
716 d_inputs.push(shadow_arg);
717 }
718 DiffActivity::Dual | DiffActivity::DualOnly => {
719 let mut shadow_arg = arg.clone();
720 let old_name = if let PatKind::Ident(_, ident, _) = arg.pat.kind {
721 ident.name
722 } else {
723 debug!("{:#?}", &shadow_arg.pat);
724 panic!("not an ident?");
725 };
726 let name: String = format!("b{}", old_name);
727 new_inputs.push(name.clone());
728 let ident = Ident::from_str_and_span(&name, shadow_arg.pat.span);
729 shadow_arg.pat = P(ast::Pat {
730 id: ast::DUMMY_NODE_ID,
731 kind: PatKind::Ident(BindingMode::NONE, ident, None),
732 span: shadow_arg.pat.span,
733 tokens: shadow_arg.pat.tokens.clone(),
734 });
735 d_inputs.push(shadow_arg);
736 }
737 DiffActivity::Const => {
738 }
740 DiffActivity::None | DiffActivity::FakeActivitySize => {
741 panic!("Should not happen");
742 }
743 }
744 if let PatKind::Ident(_, ident, _) = arg.pat.kind {
745 idents.push(ident.clone());
746 } else {
747 panic!("not an ident?");
748 }
749 }
750
751 let active_only_ret = x.ret_activity == DiffActivity::ActiveOnly;
752 if active_only_ret {
753 assert!(x.mode.is_rev());
754 }
755
756 if x.mode.is_rev() {
759 match x.ret_activity {
760 DiffActivity::Active | DiffActivity::ActiveOnly => {
761 let ty = match d_decl.output {
762 FnRetTy::Ty(ref ty) => ty.clone(),
763 FnRetTy::Default(span) => {
764 panic!("Did not expect Default ret ty: {:?}", span);
765 }
766 };
767 let name = "dret".to_string();
768 let ident = Ident::from_str_and_span(&name, ty.span);
769 let shadow_arg = ast::Param {
770 attrs: ThinVec::new(),
771 ty: ty.clone(),
772 pat: P(ast::Pat {
773 id: ast::DUMMY_NODE_ID,
774 kind: PatKind::Ident(BindingMode::NONE, ident, None),
775 span: ty.span,
776 tokens: None,
777 }),
778 id: ast::DUMMY_NODE_ID,
779 span: ty.span,
780 is_placeholder: false,
781 };
782 d_inputs.push(shadow_arg);
783 new_inputs.push(name);
784 }
785 _ => {}
786 }
787 }
788 d_decl.inputs = d_inputs.into();
789
790 if x.mode.is_fwd() {
791 if let DiffActivity::Dual = x.ret_activity {
792 let ty = match d_decl.output {
793 FnRetTy::Ty(ref ty) => ty.clone(),
794 FnRetTy::Default(span) => {
795 panic!("Did not expect Default ret ty: {:?}", span);
796 }
797 };
798 let kind = TyKind::Tup(thin_vec![ty.clone(), ty.clone()]);
801 let ty = P(rustc_ast::Ty { kind, id: ty.id, span: ty.span, tokens: None });
802 d_decl.output = FnRetTy::Ty(ty);
803 }
804 if let DiffActivity::DualOnly = x.ret_activity {
805 }
809 }
810
811 d_decl.output =
813 if active_only_ret { FnRetTy::Default(span) } else { d_decl.output.clone() };
814
815 trace!("act_ret: {:?}", act_ret);
816
817 if act_ret.len() > 0 {
821 let ret_ty = match d_decl.output {
822 FnRetTy::Ty(ref ty) => {
823 if !active_only_ret {
824 act_ret.insert(0, ty.clone());
825 }
826 let kind = TyKind::Tup(act_ret);
827 P(rustc_ast::Ty { kind, id: ty.id, span: ty.span, tokens: None })
828 }
829 FnRetTy::Default(span) => {
830 if act_ret.len() == 1 {
831 act_ret[0].clone()
832 } else {
833 let kind = TyKind::Tup(act_ret.iter().map(|arg| arg.clone()).collect());
834 P(rustc_ast::Ty { kind, id: ast::DUMMY_NODE_ID, span, tokens: None })
835 }
836 }
837 };
838 d_decl.output = FnRetTy::Ty(ret_ty);
839 }
840
841 let mut d_header = sig.header.clone();
842 if unsafe_activities {
843 d_header.safety = rustc_ast::Safety::Unsafe(span);
844 }
845 let d_sig = FnSig { header: d_header, decl: d_decl, span };
846 trace!("Generated signature: {:?}", d_sig);
847 (d_sig, new_inputs, idents, false)
848 }
849}
850
851pub(crate) use llvm_enzyme::expand;