#include #include #include #include #include #include #include #include #include #include #include #include #include #include #include import mean_field; import test_helpers; import experiment; using namespace experiment; static std::string escape_csv(const std::string& value) { if (value.find_first_of(",\"\n") == std::string::npos) { return value; } std::string escaped{"\""}; for (const char character : value) { if (character == '\"') { escaped += "\"\""; } else { escaped += character; } } escaped += '\"'; return escaped; } static void write_experiment_csv() { int rank = 0; MPI_Comm_rank(MPI_COMM_WORLD, &rank); if (rank != 0) { return; } const std::vector results = ExperimentRegistry::instance().results(); std::set parameter_names; std::set metric_names; for (const ExperimentResult& result : results) { for (const auto& [name, value] : result.parameters) { parameter_names.insert(name); } for (const auto& [name, value] : result.metrics) { metric_names.insert(name); } } const std::string output_path = ExperimentRegistry::instance().output_path(); std::ofstream output(output_path); if (!output) { std::cerr << "Unable to write experiment results to " << output_path << '\n'; return; } output << "experiment,case"; for (const std::string& name : parameter_names) { output << ',' << escape_csv(name); } for (const std::string& name : metric_names) { output << ',' << escape_csv(name); } output << '\n'; output << std::setprecision(17); for (const ExperimentResult& result : results) { output << escape_csv(result.experiment_name) << ',' << escape_csv(result.case_name); for (const std::string& name : parameter_names) { const auto iterator = result.parameters.find(name); output << ',' << (iterator == result.parameters.end() ? "" : escape_csv(iterator->second)); } for (const std::string& name : metric_names) { const auto iterator = result.metrics.find(name); output << ','; if (iterator != result.metrics.end()) { output << iterator->second; } } output << '\n'; } std::cout << "Wrote " << results.size() << " experiment rows to " << output_path << '\n'; } class ExperimentReporter final : public Catch::StreamingReporterBase { public: using StreamingReporterBase::StreamingReporterBase; static std::string getDescription() { return "Compact console reporter that writes structured experiment measurements to CSV."; } void testCaseEnded(const Catch::TestCaseStats& statistics) override { StreamingReporterBase::testCaseEnded(statistics); const bool passed = statistics.totals.assertions.allPassed(); std::cout << (passed ? "PASS " : "FAIL ") << statistics.testInfo->name << " (" << statistics.totals.assertions.passed << " assertions)\n"; } void testRunEnded(const Catch::TestRunStats& statistics) override { StreamingReporterBase::testRunEnded(statistics); write_experiment_csv(); } }; CATCH_REGISTER_REPORTER("experiment", ExperimentReporter) int main(int argc, char* argv[]) { fourdst::config::Config config; CLI::App app{"Mean Field accuracy experiments"}; app.allow_extras(); app.set_help_flag("--config-help", "Show mean-field configuration options"); fourdst::config::register_as_cli(config, app); std::string output_path{"accuracy_budget.csv"}; app.add_option("--experiment-output", output_path, "CSV path for structured measurements"); std::vector configuration_arguments{argv[0]}; std::vector catch_arguments_from_command_line; bool parsing_catch_arguments = false; for (int index = 1; index < argc; ++index) { if (std::string_view(argv[index]) == "--catch2") { parsing_catch_arguments = true; continue; } if (parsing_catch_arguments) { catch_arguments_from_command_line.emplace_back(argv[index]); } else { configuration_arguments.emplace_back(argv[index]); } } std::vector configuration_argv; configuration_argv.reserve(configuration_arguments.size()); for (const std::string& argument : configuration_arguments) { configuration_argv.push_back(argument.c_str()); } try { app.parse(static_cast(configuration_argv.size()), configuration_argv.data()); } catch (const CLI::ParseError& error) { return app.exit(error); } std::vector catch_arguments{argv[0]}; for (const std::string& argument : app.remaining()) { catch_arguments.push_back(argument); } for (const std::string& argument : catch_arguments_from_command_line) { catch_arguments.push_back(argument); } bool has_reporter = false; for (const std::string& argument : catch_arguments) { has_reporter = has_reporter || argument == "-r" || argument == "--reporter" || argument.starts_with("-r=") || argument.starts_with("--reporter="); } if (!has_reporter) { catch_arguments.emplace_back("--reporter"); catch_arguments.emplace_back("experiment"); } std::vector catch_argv; catch_argv.reserve(catch_arguments.size()); for (const std::string& argument : catch_arguments) { catch_argv.push_back(argument.c_str()); } Catch::Session session; if (const int parse_result = session.applyCommandLine(static_cast(catch_argv.size()), catch_argv.data()); parse_result != 0) { return parse_result; } mfem::Mpi::Init(argc, argv); mfem::Device device("cpu"); mean_field::utils::Args args = config.main(); if (app.count("--mesh_file") == 0) { args.mesh_file = "sandbox.smesh"; } if (app.count("--p.rtol") == 0) { args.p.rtol = 1.0e-12; } if (app.count("--p.atol") == 0) { args.p.atol = 1.0e-12; } ExperimentRegistry::instance().set_output_path(output_path); test_utils::set_args(std::move(args)); return session.run(); }