Skip to content

Commit 2085491

Browse files
Rollup merge of #156777 - sgasho:test-fix-codegen-llvm-autodiff-2, r=ZuseZ4
Add -Zautodiff_post_passes flag to limit which llvm passes to run after enzyme to make autodiff tests more robust * Add -Zautodiff_post_passes flag to limit which llvm passes to run after enzyme to make autodiff tests more robust * Change some llvm ir check in codegen_llvm/autodiff tests r? @ZuseZ4
2 parents cef63f3 + 2a9e45d commit 2085491

18 files changed

Lines changed: 209 additions & 162 deletions

File tree

compiler/rustc_codegen_llvm/src/back/lto.rs

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -620,7 +620,9 @@ pub(crate) fn run_pass_manager(
620620
if cfg!(feature = "llvm_enzyme") && enable_ad && !thin {
621621
let opt_stage = llvm::OptStage::FatLTO;
622622
let stage = write::AutodiffStage::PostAD;
623-
if !config.autodiff.contains(&config::AutoDiff::NoPostopt) {
623+
if !config.autodiff.contains(&config::AutoDiff::NoPostopt)
624+
&& config.autodiff_post_passes.as_deref() != Some("")
625+
{
624626
unsafe {
625627
write::llvm_optimize(
626628
cgcx, prof, dcx, module, None, None, config, opt_level, opt_stage, stage,

compiler/rustc_codegen_llvm/src/back/write.rs

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -566,6 +566,14 @@ pub(crate) unsafe fn llvm_optimize(
566566
let print_before_enzyme = config.autodiff.contains(&config::AutoDiff::PrintModBefore);
567567
let print_after_enzyme = config.autodiff.contains(&config::AutoDiff::PrintModAfter);
568568
let print_passes = config.autodiff.contains(&config::AutoDiff::PrintPasses);
569+
let passes_after_enzyme = if autodiff_stage == AutodiffStage::PostAD {
570+
config.autodiff_post_passes.as_deref()
571+
} else {
572+
None
573+
};
574+
let passes_after_enzyme_ptr =
575+
passes_after_enzyme.map_or(std::ptr::null(), |s| s.as_c_char_ptr());
576+
let passes_after_enzyme_len = passes_after_enzyme.map_or(0, |s| s.len());
569577
let merge_functions;
570578
let unroll_loops;
571579
let vectorize_slp;
@@ -795,6 +803,8 @@ pub(crate) unsafe fn llvm_optimize(
795803
llvm_selfprofiler,
796804
selfprofile_before_pass_callback,
797805
selfprofile_after_pass_callback,
806+
passes_after_enzyme_ptr,
807+
passes_after_enzyme_len,
798808
extra_passes.as_c_char_ptr(),
799809
extra_passes.len(),
800810
llvm_plugins.as_c_char_ptr(),

compiler/rustc_codegen_llvm/src/llvm/ffi.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2492,6 +2492,8 @@ unsafe extern "C" {
24922492
llvm_selfprofiler: *mut c_void,
24932493
begin_callback: SelfProfileBeforePassCallback,
24942494
end_callback: SelfProfileAfterPassCallback,
2495+
PostEnzymePasses: *const c_char,
2496+
PostEnzymePassesLen: size_t,
24952497
ExtraPasses: *const c_char,
24962498
ExtraPassesLen: size_t,
24972499
LLVMPlugins: *const c_char,

compiler/rustc_codegen_ssa/src/back/write.rs

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -106,6 +106,7 @@ pub struct ModuleConfig {
106106
pub emit_lifetime_markers: bool,
107107
pub llvm_plugins: Vec<String>,
108108
pub autodiff: Vec<config::AutoDiff>,
109+
pub autodiff_post_passes: Option<String>,
109110
pub offload: Vec<config::Offload>,
110111
}
111112

@@ -257,6 +258,10 @@ impl ModuleConfig {
257258
emit_lifetime_markers: sess.emit_lifetime_markers(),
258259
llvm_plugins: if_regular!(sess.opts.unstable_opts.llvm_plugins.clone(), vec![]),
259260
autodiff: if_regular!(sess.opts.unstable_opts.autodiff.clone(), vec![]),
261+
autodiff_post_passes: if_regular!(
262+
sess.opts.unstable_opts.autodiff_post_passes.clone(),
263+
None
264+
),
260265
offload: if_regular!(sess.opts.unstable_opts.offload.clone(), vec![]),
261266
}
262267
}

compiler/rustc_interface/src/tests.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -788,6 +788,7 @@ fn test_unstable_options_tracking_hash() {
788788
tracked!(annotate_moves, AnnotateMoves::Enabled(Some(1234)));
789789
tracked!(assume_incomplete_release, true);
790790
tracked!(autodiff, vec![AutoDiff::Enable, AutoDiff::NoTT]);
791+
tracked!(autodiff_post_passes, Some("function(mem2reg,instsimplify,simplifycfg)".to_string()));
791792
tracked!(binary_dep_depinfo, true);
792793
tracked!(box_noalias, false);
793794
tracked!(

compiler/rustc_llvm/llvm-wrapper/PassWrapper.cpp

Lines changed: 113 additions & 102 deletions
Original file line numberDiff line numberDiff line change
@@ -616,6 +616,7 @@ extern "C" LLVMRustResult LLVMRustOptimize(
616616
bool DebugInfoForProfiling, void *LlvmSelfProfiler,
617617
LLVMRustSelfProfileBeforePassCallback BeforePassCallback,
618618
LLVMRustSelfProfileAfterPassCallback AfterPassCallback,
619+
const char *PostEnzymePasses, size_t PostEnzymePassesLen,
619620
const char *ExtraPasses, size_t ExtraPassesLen, const char *LLVMPlugins,
620621
size_t LLVMPluginsLen) {
621622
Module *TheModule = unwrap(ModuleRef);
@@ -852,120 +853,130 @@ extern "C" LLVMRustResult LLVMRustOptimize(
852853
raw_string_ostream ThinLinkDataOS(ThinLTOSummaryBuffer->data);
853854
bool IsLTO = OptStage == LLVMRustOptStage::ThinLTO ||
854855
OptStage == LLVMRustOptStage::FatLTO;
855-
if (!NoPrepopulatePasses) {
856-
for (const auto &C : PipelineStartEPCallbacks)
857-
PB.registerPipelineStartEPCallback(C);
858-
for (const auto &C : OptimizerLastEPCallbacks)
859-
PB.registerOptimizerLastEPCallback(C);
860-
861-
// The pre-link pipelines don't support O0 and require using
862-
// buildO0DefaultPipeline() instead. At the same time, the LTO pipelines do
863-
// support O0 and using them is required.
864-
if (OptLevel == OptimizationLevel::O0 && !IsLTO) {
865-
// We manually schedule ThinLTOBufferPasses below, so don't pass the value
866-
// to enable it here.
867-
MPM = PB.buildO0DefaultPipeline(OptLevel);
868-
} else {
869-
switch (OptStage) {
870-
case LLVMRustOptStage::PreLinkNoLTO:
871-
if (ThinLTOBufferRef) {
872-
// This is similar to LLVM's `buildFatLTODefaultPipeline`, where the
873-
// bitcode for embedding is obtained after performing
874-
// `ThinLTOPreLinkDefaultPipeline`.
875-
MPM.addPass(PB.buildThinLTOPreLinkDefaultPipeline(OptLevel));
876-
MPM.addPass(ThinLTOBitcodeWriterPass(
877-
ThinLTODataOS,
878-
ThinLTOSummaryBufferRef ? &ThinLinkDataOS : nullptr));
879-
*ThinLTOBufferRef = ThinLTOBuffer.release();
880-
if (ThinLTOSummaryBufferRef) {
881-
*ThinLTOSummaryBufferRef = ThinLTOSummaryBuffer.release();
856+
if (PostEnzymePassesLen) {
857+
if (auto Err = PB.parsePassPipeline(
858+
MPM, StringRef(PostEnzymePasses, PostEnzymePassesLen))) {
859+
std::string ErrMsg = toString(std::move(Err));
860+
LLVMRustSetLastError(ErrMsg.c_str());
861+
return LLVMRustResult::Failure;
862+
}
863+
} else {
864+
if (!NoPrepopulatePasses) {
865+
for (const auto &C : PipelineStartEPCallbacks)
866+
PB.registerPipelineStartEPCallback(C);
867+
for (const auto &C : OptimizerLastEPCallbacks)
868+
PB.registerOptimizerLastEPCallback(C);
869+
870+
// The pre-link pipelines don't support O0 and require using
871+
// buildO0DefaultPipeline() instead. At the same time, the LTO pipelines
872+
// do support O0 and using them is required.
873+
if (OptLevel == OptimizationLevel::O0 && !IsLTO) {
874+
// We manually schedule ThinLTOBufferPasses below, so don't pass the
875+
// value to enable it here.
876+
MPM = PB.buildO0DefaultPipeline(OptLevel);
877+
} else {
878+
switch (OptStage) {
879+
case LLVMRustOptStage::PreLinkNoLTO:
880+
if (ThinLTOBufferRef) {
881+
// This is similar to LLVM's `buildFatLTODefaultPipeline`, where the
882+
// bitcode for embedding is obtained after performing
883+
// `ThinLTOPreLinkDefaultPipeline`.
884+
MPM.addPass(PB.buildThinLTOPreLinkDefaultPipeline(OptLevel));
885+
MPM.addPass(ThinLTOBitcodeWriterPass(
886+
ThinLTODataOS,
887+
ThinLTOSummaryBufferRef ? &ThinLinkDataOS : nullptr));
888+
*ThinLTOBufferRef = ThinLTOBuffer.release();
889+
if (ThinLTOSummaryBufferRef) {
890+
*ThinLTOSummaryBufferRef = ThinLTOSummaryBuffer.release();
891+
}
892+
MPM.addPass(PB.buildModuleOptimizationPipeline(
893+
OptLevel, ThinOrFullLTOPhase::None));
894+
MPM.addPass(
895+
createModuleToFunctionPassAdaptor(AnnotationRemarksPass()));
896+
} else {
897+
MPM = PB.buildPerModuleDefaultPipeline(OptLevel);
882898
}
883-
MPM.addPass(PB.buildModuleOptimizationPipeline(
884-
OptLevel, ThinOrFullLTOPhase::None));
885-
MPM.addPass(
886-
createModuleToFunctionPassAdaptor(AnnotationRemarksPass()));
887-
} else {
888-
MPM = PB.buildPerModuleDefaultPipeline(OptLevel);
899+
break;
900+
case LLVMRustOptStage::PreLinkThinLTO:
901+
case LLVMRustOptStage::PreLinkFatLTO:
902+
MPM = PB.buildThinLTOPreLinkDefaultPipeline(OptLevel);
903+
NeedThinLTOBufferPasses = false;
904+
break;
905+
case LLVMRustOptStage::ThinLTO:
906+
// FIXME: Does it make sense to pass the ModuleSummaryIndex?
907+
// It only seems to be needed for C++ specific optimizations.
908+
MPM = PB.buildThinLTODefaultPipeline(OptLevel, nullptr);
909+
break;
910+
case LLVMRustOptStage::FatLTO:
911+
MPM = PB.buildLTODefaultPipeline(OptLevel, nullptr);
912+
NeedThinLTOBufferPasses = false;
913+
break;
889914
}
890-
break;
891-
case LLVMRustOptStage::PreLinkThinLTO:
892-
case LLVMRustOptStage::PreLinkFatLTO:
893-
MPM = PB.buildThinLTOPreLinkDefaultPipeline(OptLevel);
894-
NeedThinLTOBufferPasses = false;
895-
break;
896-
case LLVMRustOptStage::ThinLTO:
897-
// FIXME: Does it make sense to pass the ModuleSummaryIndex?
898-
// It only seems to be needed for C++ specific optimizations.
899-
MPM = PB.buildThinLTODefaultPipeline(OptLevel, nullptr);
900-
break;
901-
case LLVMRustOptStage::FatLTO:
902-
MPM = PB.buildLTODefaultPipeline(OptLevel, nullptr);
903-
NeedThinLTOBufferPasses = false;
904-
break;
905915
}
916+
} else {
917+
// We're not building any of the default pipelines but we still want to
918+
// add the verifier, instrumentation, etc passes if they were requested
919+
for (const auto &C : PipelineStartEPCallbacks)
920+
C(MPM, OptLevel);
921+
for (const auto &C : OptimizerLastEPCallbacks)
922+
C(MPM, OptLevel, ThinOrFullLTOPhase::None);
906923
}
907-
} else {
908-
// We're not building any of the default pipelines but we still want to
909-
// add the verifier, instrumentation, etc passes if they were requested
910-
for (const auto &C : PipelineStartEPCallbacks)
911-
C(MPM, OptLevel);
912-
for (const auto &C : OptimizerLastEPCallbacks)
913-
C(MPM, OptLevel, ThinOrFullLTOPhase::None);
914-
}
915924

916-
if (ExtraPassesLen) {
917-
if (auto Err =
918-
PB.parsePassPipeline(MPM, StringRef(ExtraPasses, ExtraPassesLen))) {
919-
std::string ErrMsg = toString(std::move(Err));
920-
LLVMRustSetLastError(ErrMsg.c_str());
921-
return LLVMRustResult::Failure;
925+
if (ExtraPassesLen) {
926+
if (auto Err = PB.parsePassPipeline(
927+
MPM, StringRef(ExtraPasses, ExtraPassesLen))) {
928+
std::string ErrMsg = toString(std::move(Err));
929+
LLVMRustSetLastError(ErrMsg.c_str());
930+
return LLVMRustResult::Failure;
931+
}
922932
}
923-
}
924933

925-
if (NeedThinLTOBufferPasses) {
926-
MPM.addPass(CanonicalizeAliasesPass());
927-
MPM.addPass(NameAnonGlobalPass());
928-
}
929-
// For `-Copt-level=0`, and the pre-link fat/thin LTO stages.
930-
if (ThinLTOBufferRef && *ThinLTOBufferRef == nullptr) {
931-
// thin lto summaries prevent fat lto, so do not emit them if fat
932-
// lto is requested. See PR #136840 for background information.
933-
if (OptStage != LLVMRustOptStage::PreLinkFatLTO) {
934-
MPM.addPass(ThinLTOBitcodeWriterPass(
935-
ThinLTODataOS, ThinLTOSummaryBufferRef ? &ThinLinkDataOS : nullptr));
936-
} else {
937-
MPM.addPass(BitcodeWriterPass(ThinLTODataOS));
934+
if (NeedThinLTOBufferPasses) {
935+
MPM.addPass(CanonicalizeAliasesPass());
936+
MPM.addPass(NameAnonGlobalPass());
938937
}
939-
*ThinLTOBufferRef = ThinLTOBuffer.release();
940-
if (ThinLTOSummaryBufferRef) {
941-
*ThinLTOSummaryBufferRef = ThinLTOSummaryBuffer.release();
938+
// For `-Copt-level=0`, and the pre-link fat/thin LTO stages.
939+
if (ThinLTOBufferRef && *ThinLTOBufferRef == nullptr) {
940+
// thin lto summaries prevent fat lto, so do not emit them if fat
941+
// lto is requested. See PR #136840 for background information.
942+
if (OptStage != LLVMRustOptStage::PreLinkFatLTO) {
943+
MPM.addPass(ThinLTOBitcodeWriterPass(
944+
ThinLTODataOS,
945+
ThinLTOSummaryBufferRef ? &ThinLinkDataOS : nullptr));
946+
} else {
947+
MPM.addPass(BitcodeWriterPass(ThinLTODataOS));
948+
}
949+
*ThinLTOBufferRef = ThinLTOBuffer.release();
950+
if (ThinLTOSummaryBufferRef) {
951+
*ThinLTOSummaryBufferRef = ThinLTOSummaryBuffer.release();
952+
}
942953
}
943-
}
944954

945-
// now load "-enzyme" pass:
946-
// With dlopen, ENZYME macro may not be defined, so check EnzymePtr directly
947-
// In the case of debug builds with multiple codegen units, we might not
948-
// have all function definitions available during the early compiler
949-
// invocations. We therefore wait for the final lto step to run Enzyme.
950-
if (EnzymePtr && IsLTO) {
951-
952-
if (PrintBeforeEnzyme) {
953-
// Handle the Rust flag `-Zautodiff=PrintModBefore`.
954-
std::string Banner = "Module before EnzymeNewPM";
955-
MPM.addPass(PrintModulePass(outs(), Banner, true, false));
956-
}
955+
// now load "-enzyme" pass:
956+
// With dlopen, ENZYME macro may not be defined, so check EnzymePtr directly
957+
// In the case of debug builds with multiple codegen units, we might not
958+
// have all function definitions available during the early compiler
959+
// invocations. We therefore wait for the final lto step to run Enzyme.
960+
if (EnzymePtr && IsLTO) {
961+
962+
if (PrintBeforeEnzyme) {
963+
// Handle the Rust flag `-Zautodiff=PrintModBefore`.
964+
std::string Banner = "Module before EnzymeNewPM";
965+
MPM.addPass(PrintModulePass(outs(), Banner, true, false));
966+
}
957967

958-
EnzymePtr(PB, false);
959-
if (auto Err = PB.parsePassPipeline(MPM, "enzyme")) {
960-
std::string ErrMsg = toString(std::move(Err));
961-
LLVMRustSetLastError(ErrMsg.c_str());
962-
return LLVMRustResult::Failure;
963-
}
968+
EnzymePtr(PB, false);
969+
if (auto Err = PB.parsePassPipeline(MPM, "enzyme")) {
970+
std::string ErrMsg = toString(std::move(Err));
971+
LLVMRustSetLastError(ErrMsg.c_str());
972+
return LLVMRustResult::Failure;
973+
}
964974

965-
if (PrintAfterEnzyme) {
966-
// Handle the Rust flag `-Zautodiff=PrintModAfter`.
967-
std::string Banner = "Module after EnzymeNewPM";
968-
MPM.addPass(PrintModulePass(outs(), Banner, true, false));
975+
if (PrintAfterEnzyme) {
976+
// Handle the Rust flag `-Zautodiff=PrintModAfter`.
977+
std::string Banner = "Module after EnzymeNewPM";
978+
MPM.addPass(PrintModulePass(outs(), Banner, true, false));
979+
}
969980
}
970981
}
971982

compiler/rustc_session/src/options.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2298,6 +2298,8 @@ options! {
22982298
`=LooseTypes`
22992299
`=Inline`
23002300
Multiple options can be combined with commas."),
2301+
autodiff_post_passes: Option<String> = (None, parse_opt_string, [TRACKED],
2302+
"set llvm passes to run after enzyme (no passes run when it is empty)"),
23012303
#[rustc_lint_opt_deny_field_access("use `Session::binary_dep_depinfo` instead of this field")]
23022304
binary_dep_depinfo: bool = (false, parse_bool, [TRACKED],
23032305
"include artifacts (sysroot, crate dependencies) used during compilation in dep-info \

0 commit comments

Comments
 (0)