feat(mean_field): added initial implementation

note this implementation lacks many tests
This commit is contained in:
2026-07-15 09:44:43 -04:00
commit 9bc4f2758a
49 changed files with 171811 additions and 0 deletions

331
tests/quadrature/policy.cpp Normal file
View File

@@ -0,0 +1,331 @@
#include <catch2/catch_test_macros.hpp>
#include <catch2/matchers/catch_matchers_floating_point.hpp>
#include <array>
#include <cmath>
#include <stdexcept>
#include <string_view>
#include <utility>
#include <mfem.hpp>
import mean_field;
import test_helpers;
using namespace mean_field;
namespace {
std::string_view get_term_name(const quadrature::Term term) {
switch (term) {
case quadrature::Term::gravity_hdiv_mass: return "gravity_hdiv_mass";
case quadrature::Term::gravity_divergence: return "gravity_divergence";
case quadrature::Term::gravity_source: return "gravity_source";
case quadrature::Term::gravity_boundary: return "gravity_boundary";
case quadrature::Term::density_projection: return "density_projection";
case quadrature::Term::mass_conservation: return "mass_conservation";
case quadrature::Term::center_of_mass: return "center_of_mass";
case quadrature::Term::quadrupole: return "quadrupole";
case quadrature::Term::gravitational_energy: return "gravitational_energy";
case quadrature::Term::virial: return "virial";
case quadrature::Term::error_norm: return "error_norm";
}
return "unknown";
}
quadrature::Query make_query(const quadrature::Term term, const int base_order) {
return {.term = term, .base_order = base_order};
}
}
TEST_CASE("Quadrature Policy Computes Base Orders", tags::unit & tags::quadrature) {
const quadrature::Policy policy(quadrature::make_rule_set(quadrature::Mode::production));
quadrature::Query generic_query = {
.term = quadrature::Term::gravitational_energy,
.trial_order = 3,
.test_order = 4,
.coefficient_order = 2,
.geometry_weight_order = 5
};
const quadrature::Resolution generic_resolution = policy.resolve(generic_query);
CHECK(generic_resolution.base_order == 14);
CHECK(generic_resolution.boost == 0);
CHECK(generic_resolution.order == 14);
CHECK_FALSE(generic_resolution.used_fixed_order);
quadrature::Query divergence_query = {
.term = quadrature::Term::gravity_divergence,
.trial_order = 3,
.test_order = 2,
.coefficient_order = 1,
.geometry_weight_order = 4
};
CHECK(policy.resolve(divergence_query).base_order == 9);
divergence_query.trial_order = 0;
CHECK(policy.resolve(divergence_query).base_order == 7);
quadrature::Query explicit_query = {
.term = quadrature::Term::gravity_hdiv_mass,
.trial_order = 20,
.test_order = 20,
.coefficient_order = 20,
.geometry_weight_order = 20,
.base_order = 11
};
CHECK(policy.resolve(explicit_query).base_order == 11);
CHECK(policy.resolve(explicit_query).order == 11);
}
TEST_CASE("Quadrature Policy Composes Global and Term Boosts", tags::unit & tags::quadrature) {
quadrature::RuleSet rule_set = quadrature::make_rule_set(quadrature::Mode::production, 3);
rule_set.gravity_hdiv_mass.boost = 5;
rule_set.error_norm.boost = 2;
const quadrature::Policy policy(rule_set);
const quadrature::Resolution mass_resolution = policy.resolve(make_query(quadrature::Term::gravity_hdiv_mass, 7));
CHECK(mass_resolution.base_order == 7);
CHECK(mass_resolution.boost == 8);
CHECK(mass_resolution.order == 15);
CHECK_FALSE(mass_resolution.used_fixed_order);
const quadrature::Resolution error_resolution = policy.resolve(make_query(quadrature::Term::error_norm, 7));
CHECK(error_resolution.boost == 5);
CHECK(error_resolution.order == 12);
const quadrature::Resolution source_resolution = policy.resolve(make_query(quadrature::Term::gravity_source, 7));
CHECK(source_resolution.boost == 3);
CHECK(source_resolution.order == 10);
}
TEST_CASE("Quadrature Fixed Orders Have Defined Precedence", tags::unit & tags::quadrature) {
quadrature::RuleSet rule_set = quadrature::make_rule_set(quadrature::Mode::production, 4);
rule_set.fallback.fixed_order = 17;
rule_set.gravity_hdiv_mass.fixed_order = 23;
rule_set.gravity_hdiv_mass.boost = 100;
rule_set.gravity_source.boost = 100;
const quadrature::Policy policy(rule_set);
const quadrature::Resolution term_resolution = policy.resolve(make_query(quadrature::Term::gravity_hdiv_mass, 8));
CHECK(term_resolution.base_order == 8);
CHECK(term_resolution.boost == 0);
CHECK(term_resolution.order == 23);
CHECK(term_resolution.used_fixed_order);
const quadrature::Resolution fallback_resolution = policy.resolve(make_query(quadrature::Term::gravity_source, 8));
CHECK(fallback_resolution.base_order == 8);
CHECK(fallback_resolution.boost == 0);
CHECK(fallback_resolution.order == 17);
CHECK(fallback_resolution.used_fixed_order);
}
TEST_CASE("Quadrature Modes Apply Their Expected Baseline Boosts", tags::unit & tags::quadrature) {
constexpr int base_order = 6;
constexpr int global_boost = 3;
for (const quadrature::Mode mode : {quadrature::Mode::fast, quadrature::Mode::production, quadrature::Mode::convergence}) {
const quadrature::Policy policy(quadrature::make_rule_set(mode, global_boost));
const quadrature::Resolution resolution = policy.resolve(make_query(quadrature::Term::error_norm, base_order));
CHECK(resolution.boost == global_boost);
CHECK(resolution.order == base_order + global_boost);
}
const quadrature::Policy reference_policy(quadrature::make_rule_set(quadrature::Mode::reference, global_boost));
const quadrature::Resolution reference_resolution = reference_policy.resolve(make_query(quadrature::Term::error_norm, base_order));
CHECK(reference_resolution.boost == global_boost + 8);
CHECK(reference_resolution.order == base_order + global_boost + 8);
}
TEST_CASE("Quadrature Policy Routes Every Term to Its Control", tags::unit & tags::quadrature) {
quadrature::RuleSet rule_set;
rule_set.gravity_hdiv_mass.boost = 1;
rule_set.gravity_divergence.boost = 2;
rule_set.gravity_source.boost = 3;
rule_set.gravity_boundary.boost = 4;
rule_set.density_projection.boost = 5;
rule_set.mass_conservation.boost = 6;
rule_set.center_of_mass.boost = 7;
rule_set.quadrupole.boost = 8;
rule_set.gravitational_energy.boost = 9;
rule_set.virial.boost = 10;
rule_set.error_norm.boost = 11;
const quadrature::Policy policy(rule_set);
const std::array<std::pair<quadrature::Term, int>, 11> cases = {{
{quadrature::Term::gravity_hdiv_mass, 1},
{quadrature::Term::gravity_divergence, 2},
{quadrature::Term::gravity_source, 3},
{quadrature::Term::gravity_boundary, 4},
{quadrature::Term::density_projection, 5},
{quadrature::Term::mass_conservation, 6},
{quadrature::Term::center_of_mass, 7},
{quadrature::Term::quadrupole, 8},
{quadrature::Term::gravitational_energy, 9},
{quadrature::Term::virial, 10},
{quadrature::Term::error_norm, 11}
}};
for (const auto& [term, expected_boost] : cases) {
DYNAMIC_SECTION(get_term_name(term)) {
const quadrature::Resolution resolution = policy.resolve(make_query(term, 20));
CHECK(resolution.boost == expected_boost);
CHECK(resolution.order == 20 + expected_boost);
}
}
}
TEST_CASE("Quadrature Policy Rejects Invalid Orders", tags::unit & tags::quadrature) {
const quadrature::Policy policy(quadrature::make_rule_set(quadrature::Mode::production));
quadrature::Query negative_component_query = {
.term = quadrature::Term::error_norm,
.trial_order = -1
};
REQUIRE_THROWS_AS(policy.resolve(negative_component_query), std::invalid_argument);
quadrature::Query negative_base_query = {
.term = quadrature::Term::error_norm,
.base_order = -1
};
REQUIRE_THROWS_AS(policy.resolve(negative_base_query), std::invalid_argument);
quadrature::RuleSet negative_fixed_rule_set;
negative_fixed_rule_set.error_norm.fixed_order = -1;
const quadrature::Policy negative_fixed_policy(negative_fixed_rule_set);
REQUIRE_THROWS_AS(negative_fixed_policy.resolve(make_query(quadrature::Term::error_norm, 3)), std::invalid_argument);
quadrature::RuleSet negative_resolved_rule_set;
negative_resolved_rule_set.fallback.boost = -4;
const quadrature::Policy negative_resolved_policy(negative_resolved_rule_set);
REQUIRE_THROWS_AS(negative_resolved_policy.resolve(make_query(quadrature::Term::error_norm, 3)), std::invalid_argument);
}
TEST_CASE("MFEM Rule Factory Returns the Resolved Rule", tags::unit & tags::quadrature) {
quadrature::RuleSet rule_set = quadrature::make_rule_set(quadrature::Mode::production, 2);
rule_set.error_norm.boost = 3;
const quadrature::RuleFactory factory{quadrature::Policy(rule_set)};
const quadrature::MfemRule selected_rule = factory.get(make_query(quadrature::Term::error_norm, 4), mfem::Geometry::CUBE);
const mfem::IntegrationRule& expected_rule = mfem::IntRules.Get(mfem::Geometry::CUBE, 9);
REQUIRE(selected_rule.integration_rule != nullptr);
CHECK(selected_rule.resolution.base_order == 4);
CHECK(selected_rule.resolution.boost == 5);
CHECK(selected_rule.resolution.order == 9);
CHECK(selected_rule.integration_rule == &expected_rule);
CHECK(selected_rule.integration_rule->GetNPoints() > 0);
}
TEST_CASE("MFEM Quadrature Rule Integrates Tensor Polynomial Exactly", tags::unit & tags::quadrature) {
constexpr int polynomial_degree = 7;
const quadrature::RuleFactory factory{quadrature::Policy(quadrature::make_rule_set(quadrature::Mode::production))};
const quadrature::MfemRule selected_rule = factory.get(make_query(quadrature::Term::error_norm, polynomial_degree), mfem::Geometry::CUBE);
double numerical_integral = 0.0;
for (int i = 0; i < selected_rule.integration_rule->GetNPoints(); ++i) {
const mfem::IntegrationPoint& integration_point = selected_rule.integration_rule->IntPoint(i);
numerical_integral += integration_point.weight * std::pow(integration_point.x, polynomial_degree) * std::pow(integration_point.y, polynomial_degree) * std::pow(integration_point.z, polynomial_degree);
}
const double one_dimensional_integral = 1.0 / static_cast<double>(polynomial_degree + 1);
const double analytic_integral = one_dimensional_integral * one_dimensional_integral * one_dimensional_integral;
CHECK_THAT(numerical_integral, Catch::Matchers::WithinAbs(analytic_integral, 5.0e-14));
}
TEST_CASE("Policy Controlled Hdiv Mass Assembly Matches Overintegrated Reference", tags::quadrature & tags::solver & tags::integration) {
mfem::Mesh mesh = mfem::Mesh::MakeCartesian3D(1, 1, 1, mfem::Element::HEXAHEDRON, 1.0, 1.0, 1.0);
mfem::RT_FECollection rt_collection(2, 3);
mfem::FiniteElementSpace rt_space(&mesh, &rt_collection);
const mfem::FiniteElement* rt_element = rt_space.GetTypicalFE();
mfem::ElementTransformation* transformation = mesh.GetElementTransformation(0);
const int base_order = 2 * rt_element->GetOrder() + transformation->OrderW();
const quadrature::RuleFactory production_factory{quadrature::Policy(quadrature::make_rule_set(quadrature::Mode::production))};
const quadrature::MfemRule production_rule = production_factory.get(make_query(quadrature::Term::gravity_hdiv_mass, base_order), rt_element->GetGeomType());
quadrature::RuleSet reference_rule_set = quadrature::make_rule_set(quadrature::Mode::production);
reference_rule_set.gravity_hdiv_mass.boost = 8;
const quadrature::RuleFactory reference_factory{quadrature::Policy(reference_rule_set)};
const quadrature::MfemRule reference_rule = reference_factory.get(make_query(quadrature::Term::gravity_hdiv_mass, base_order), rt_element->GetGeomType());
mfem::BilinearForm production_mass(&rt_space);
auto* production_integrator = new mfem::VectorFEMassIntegrator();
production_integrator->SetIntegrationRule(*production_rule.integration_rule);
production_mass.AddDomainIntegrator(production_integrator);
production_mass.Assemble();
production_mass.Finalize();
mfem::BilinearForm reference_mass(&rt_space);
auto* reference_integrator = new mfem::VectorFEMassIntegrator();
reference_integrator->SetIntegrationRule(*reference_rule.integration_rule);
reference_mass.AddDomainIntegrator(reference_integrator);
reference_mass.Assemble();
reference_mass.Finalize();
mfem::Vector input(rt_space.GetVSize());
mfem::Vector production_output(rt_space.GetVSize());
mfem::Vector reference_output(rt_space.GetVSize());
for (int i = 0; i < input.Size(); ++i) {
input(i) = std::sin(0.37 * static_cast<double>(i + 1));
}
production_mass.Mult(input, production_output);
reference_mass.Mult(input, reference_output);
mfem::Vector difference(production_output);
difference -= reference_output;
const double relative_difference = difference.Norml2() / reference_output.Norml2();
INFO("Production quadrature order = " << production_rule.resolution.order);
INFO("Reference quadrature order = " << reference_rule.resolution.order);
INFO("Relative operator difference = " << relative_difference);
CHECK_THAT(relative_difference, Catch::Matchers::WithinAbs(0.0, 1.0e-12));
}
TEST_CASE("HDiv Mass Helper Resolves the MFEM Baseline", tags::unit & tags::quadrature & tags::solver) {
mfem::Mesh mesh = mfem::Mesh::MakeCartesian3D(1, 1, 1, mfem::Element::HEXAHEDRON);
mfem::RT_FECollection rt_collection(2, 3);
mfem::FiniteElementSpace rt_space(&mesh, &rt_collection);
quadrature::RuleSet rule_set = quadrature::make_rule_set(quadrature::Mode::production);
rule_set.gravity_hdiv_mass.boost = 3;
quadrature::RuleFactory factory{quadrature::Policy(std::move(rule_set))};
const mfem::FiniteElement& element = *rt_space.GetTypicalFE();
const mfem::ElementTransformation& transformation = *mesh.GetElementTransformation(0);
mfem::VectorFEMassIntegrator integrator;
const quadrature::Resolution resolution = factory.configure_gravity_hdiv_mass(integrator, quadrature::QuadratureRole::discretization, element, transformation);
const int expected_base_order = 2 * element.GetOrder() + transformation.OrderW();
CHECK(resolution.base_order == expected_base_order);
CHECK(resolution.boost == 3);
CHECK(resolution.order == expected_base_order + 3);
}
TEST_CASE("Gravity Divergence Helper Resolves Preconditioner Rule", tags::unit & tags::quadrature & tags::solver) {
mfem::Mesh mesh = mfem::Mesh::MakeCartesian3D(1, 1, 1, mfem::Element::HEXAHEDRON);
mfem::RT_FECollection rt_collection(2, 3);
mfem::L2_FECollection l2_collection(2, 3);
mfem::FiniteElementSpace rt_space(&mesh, &rt_collection);
mfem::FiniteElementSpace l2_space(&mesh, &l2_collection);
quadrature::RuleSet rule_set = quadrature::make_rule_set(quadrature::Mode::production);
rule_set.gravity_divergence.boost = 2;
rule_set.roles.preconditioner.boost = 3;
quadrature::RuleFactory factory{quadrature::Policy(std::move(rule_set))};
const mfem::FiniteElement& trial_element = *rt_space.GetTypicalFE();
const mfem::FiniteElement& test_element = *l2_space.GetTypicalFE();
const mfem::ElementTransformation& transformation = *mesh.GetElementTransformation(0);
mfem::VectorFEDivergenceIntegrator integrator;
const quadrature::Resolution resolution = factory.configure_gravity_divergence(integrator, quadrature::QuadratureRole::preconditioner, trial_element, test_element, transformation);
const int expected_base_order = std::max(0, trial_element.GetOrder() - 1) + test_element.GetOrder() + transformation.OrderW();
CHECK(resolution.base_order == expected_base_order);
CHECK(resolution.boost == 5);
CHECK(resolution.order == expected_base_order + 5);
}