restricted the unknown state vector to surface deformation and implemented one prescription, NodalRadialSurface, while the full volumetric displacment field is reconstructed analytically from that. This reduced the number of degrees of freedom in the system by a factor of 80 while also removing many null vectors from the system.
215 lines
6.6 KiB
C++
215 lines
6.6 KiB
C++
#include <catch2/catch_session.hpp>
|
|
#include <catch2/reporters/catch_reporter_registrars.hpp>
|
|
#include <catch2/reporters/catch_reporter_streaming_base.hpp>
|
|
|
|
#include <CLI/CLI.hpp>
|
|
#include <fourdst/config/config.h>
|
|
#include <mfem.hpp>
|
|
|
|
#include <catch2/catch_test_case_info.hpp>
|
|
#include <fstream>
|
|
#include <iomanip>
|
|
#include <iostream>
|
|
#include <set>
|
|
#include <sstream>
|
|
#include <string>
|
|
#include <string_view>
|
|
#include <vector>
|
|
|
|
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<ExperimentResult> results = ExperimentRegistry::instance().results();
|
|
std::set<std::string> parameter_names;
|
|
std::set<std::string> 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<mean_field::utils::Args> 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<std::string> configuration_arguments{argv[0]};
|
|
std::vector<std::string> 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<const char *> 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<int>(configuration_argv.size()), configuration_argv.data());
|
|
} catch (const CLI::ParseError &error) {
|
|
return app.exit(error);
|
|
}
|
|
|
|
std::vector<std::string> 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<const char *> 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<int>(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();
|
|
}
|