perf(jacobian-action): major updates to jacobian action application by removing redudant quadrature work. ~5x increase in speed

This commit is contained in:
2026-09-02 17:01:50 -04:00
parent 85500fef3b
commit 25510008dd
74 changed files with 8967 additions and 814 deletions

View File

@@ -1,5 +1,6 @@
module;
#include <array>
#include <mfem.hpp>
module mean_field;
@@ -7,6 +8,75 @@ module mean_field;
import :operators.kernels.rotational_displacement_force;
import :operators.prepared_rotational_displacement_force;
namespace {
using DomainSchema = mean_field::utils::domain::CoreEnvelopeVacuumDomainSchema;
[[nodiscard]] bool is_vacuum_attribute(const int attribute) {
return DomainSchema::template attribute_belongs_to<mean_field::utils::domain::Vacuum>(attribute);
}
void true_to_local(
const mfem::ParFiniteElementSpace &finiteElementSpace,
const mfem::Vector &trueVector,
mfem::Vector &localVector
) {
localVector.SetSize(finiteElementSpace.GetVSize());
const mfem::Operator *prolongation = finiteElementSpace.GetProlongationMatrix();
if (prolongation != nullptr) {
prolongation->Mult(trueVector, localVector);
} else {
localVector = trueVector;
}
}
void local_to_true(
const mfem::ParFiniteElementSpace &finiteElementSpace,
const mfem::Vector &localVector,
mfem::Vector &trueVector
) {
trueVector.SetSize(finiteElementSpace.GetTrueVSize());
trueVector = 0.0;
const mfem::Operator *prolongation = finiteElementSpace.GetProlongationMatrix();
if (prolongation != nullptr) {
prolongation->MultTranspose(localVector, trueVector);
} else {
trueVector = localVector;
}
}
[[nodiscard]] int vector_dof_index(
const mfem::Ordering::Type ordering,
const int scalarDof,
const int component,
const int scalarDofCount,
const int dimension
) {
if (ordering == mfem::Ordering::byNODES) {
return scalarDof + component * scalarDofCount;
}
MFEM_VERIFY(ordering == mfem::Ordering::byVDIM, "Unsupported displacement ordering.");
return scalarDof * dimension + component;
}
[[nodiscard]] const mfem::IntegrationRule &get_rotation_force_rule(
const mean_field::fem::FEM &f,
const mfem::ElementTransformation &transformation
) {
using DisplacementField = mean_field::field::Field<mean_field::field::Displacement>;
const mean_field::quadrature::Query query =
DisplacementField::make_query<mean_field::field::Displacement::Form::CentrifugalForce>(
mean_field::quadrature::QuadratureRole::discretization, transformation.OrderW(), std::array<int, 1>{1},
mean_field::utils::DOMAINS::STELLAR, mean_field::quadrature::MappingKind::general
);
const mean_field::quadrature::MfemRule rule = f.quadratureFactory->get(query, transformation.GetGeometryType());
MFEM_VERIFY(
rule.integration_rule != nullptr,
"The quadrature policy did not return a rotational-displacement-force integration rule."
);
return *rule.integration_rule;
}
} // namespace
namespace mean_field::operators {
PreparedRotationalDisplacementForceOperator::PreparedRotationalDisplacementForceOperator(
const fem::FEM &f,
@@ -49,6 +119,105 @@ namespace mean_field::operators {
);
}
void PreparedRotationalDisplacementForceOperator::PrepareElementData() {
MFEM_VERIFY(m_rotation.has_value(), "Prepared rotational force has no frozen rotation state.");
m_elements.clear();
m_elements.reserve(m_fem.mesh->GetNE());
mfem::Vector baseDensityLocal;
mfem::Vector baseDisplacementLocal;
true_to_local(*m_fem.densityFes, m_context.GetBaseDensityTrue(), baseDensityLocal);
true_to_local(*m_fem.displacementFes, m_context.GetDisplacementTrue(), baseDisplacementLocal);
mapping::DomainMapper::Workspace workspace(m_domainMapper.GetDimension());
mapping::VolumeMappingContext mappingContext;
mfem::Array<int> compactificationDofs;
mfem::Vector elementBaseDensity;
mfem::Vector elementBaseDisplacement;
mfem::Vector elementCompactification;
mfem::Vector densityShape;
mfem::Vector potentialGradient;
const int dimension = m_domainMapper.GetDimension();
for (int elementId = 0; elementId < m_fem.mesh->GetNE(); ++elementId) {
mfem::ElementTransformation *transformation = m_fem.mesh->GetElementTransformation(elementId);
MFEM_VERIFY(transformation != nullptr, "Prepared rotational force received a null transformation.");
if (is_vacuum_attribute(transformation->Attribute)) {
continue;
}
m_elements.emplace_back();
ElementPAData &data = m_elements.back();
data.elementId = elementId;
data.densityDofTransformation = m_fem.densityFes->GetElementDofs(elementId, data.densityDofs);
data.displacementDofTransformation =
m_fem.displacementFes->GetElementVDofs(elementId, data.displacementDofs);
mfem::DofTransformation *compactificationDofTransformation =
m_fem.compactificationFes->GetElementDofs(elementId, compactificationDofs);
baseDensityLocal.GetSubVector(data.densityDofs, elementBaseDensity);
baseDisplacementLocal.GetSubVector(data.displacementDofs, elementBaseDisplacement);
m_fem.compactificationCoordinate->GetSubVector(compactificationDofs, elementCompactification);
if (data.densityDofTransformation != nullptr) {
data.densityDofTransformation->InvTransformPrimal(elementBaseDensity);
}
if (data.displacementDofTransformation != nullptr) {
data.displacementDofTransformation->InvTransformPrimal(elementBaseDisplacement);
}
if (compactificationDofTransformation != nullptr) {
compactificationDofTransformation->InvTransformPrimal(elementCompactification);
}
const mfem::FiniteElement &densityElement = *m_fem.densityFes->GetFE(elementId);
const mfem::FiniteElement &displacementElement = *m_fem.displacementFes->GetFE(elementId);
const mfem::FiniteElement &compactificationElement = *m_fem.compactificationFes->GetFE(elementId);
data.integrationRule = &get_rotation_force_rule(m_fem, *transformation);
const mapping::ElementDisplacementData displacementData =
mapping::ElementDisplacementDataFromElementVDofs(displacementElement, elementBaseDisplacement);
const mapping::ElementCompactificationData compactificationData(
compactificationElement, elementCompactification
);
const mapping::ElementMappingData mappingData{
.displacement = displacementData, .compactification = compactificationData
};
const int quadraturePointCount = data.integrationRule->GetNPoints();
data.inverseElementJacobians.SetSize(quadraturePointCount, dimension * dimension);
data.centrifugalAccelerations.SetSize(quadraturePointCount, dimension);
data.baseDensityValues.SetSize(quadraturePointCount);
data.quadratureWeights.SetSize(quadraturePointCount);
densityShape.SetSize(densityElement.GetDof());
potentialGradient.SetSize(dimension);
for (int quadraturePoint = 0; quadraturePoint < quadraturePointCount; ++quadraturePoint) {
const mfem::IntegrationPoint &integrationPoint = data.integrationRule->IntPoint(quadraturePoint);
const mapping::MappingStatus status = m_domainMapper.EvaluateVolume(
mappingData, *transformation, integrationPoint, workspace, mappingContext
);
MFEM_VERIFY(
status == mapping::MappingStatus::valid && !mappingContext.mapping.compactified,
"Prepared rotational force encountered an invalid stellar mapping."
);
densityElement.CalcShape(integrationPoint, densityShape);
m_rotation->potential_gradient(mappingContext.mapping.physical_position, potentialGradient);
data.baseDensityValues(quadraturePoint) = elementBaseDensity * densityShape;
data.quadratureWeights(quadraturePoint) = mappingContext.quadrature.weight;
for (int row = 0; row < dimension; ++row) {
data.centrifugalAccelerations(quadraturePoint, row) = -potentialGradient(row);
for (int column = 0; column < dimension; ++column) {
data.inverseElementJacobians(quadraturePoint, row * dimension + column) =
mappingContext.quadrature.J_inv(row, column);
}
}
}
}
}
PreparedRotationalDisplacementForceReport PreparedRotationalDisplacementForceOperator::Prepare(
const context::rotational_displacement_force::RotationalDisplacementForceStateView &state,
const context::rotational_displacement_force::RotationalDisplacementForceDependencies &dependencies,
@@ -83,6 +252,7 @@ namespace mean_field::operators {
);
m_cachedResidual.SetSize(m_context.GetDisplacementMap().reduced_size());
m_context.GetDisplacementMap().gather(m_actionTrue, m_cachedResidual);
PrepareElementData();
++m_residualPreparationCount;
report.preparedResidual = true;
@@ -143,6 +313,97 @@ namespace mean_field::operators {
++m_displacementJacobianStatistics.applications;
}
void PreparedRotationalDisplacementForceOperator::ApplyPreparedCompleteJacobianActionTrue(
const mfem::Vector &densityVariationTrue,
const mfem::Vector &displacementVariationTrue,
mfem::Vector &actionTrue
) const {
true_to_local(*m_fem.densityFes, densityVariationTrue, m_densityVariationLocal);
true_to_local(*m_fem.displacementFes, displacementVariationTrue, m_displacementVariationLocal);
m_localAction.SetSize(m_fem.displacementFes->GetVSize());
m_localAction = 0.0;
const int dimension = m_domainMapper.GetDimension();
const mfem::Ordering::Type ordering = m_fem.displacementFes->GetOrdering();
for (const ElementPAData &data : m_elements) {
MFEM_VERIFY(data.integrationRule != nullptr, "Prepared rotational force has no integration rule.");
m_densityVariationLocal.GetSubVector(data.densityDofs, m_elementDensityVariation);
m_displacementVariationLocal.GetSubVector(data.displacementDofs, m_elementDisplacementVariation);
if (data.densityDofTransformation != nullptr) {
data.densityDofTransformation->InvTransformPrimal(m_elementDensityVariation);
}
if (data.displacementDofTransformation != nullptr) {
data.displacementDofTransformation->InvTransformPrimal(m_elementDisplacementVariation);
}
const mfem::FiniteElement &densityElement = *m_fem.densityFes->GetFE(data.elementId);
const mfem::FiniteElement &displacementElement = *m_fem.displacementFes->GetFE(data.elementId);
const mapping::ElementDisplacementData directionData =
mapping::ElementDisplacementDataFromElementVDofs(displacementElement, m_elementDisplacementVariation);
const mfem::DenseMatrix &directionDofs = directionData.GetDofMatrix();
const int scalarDisplacementDofCount = displacementElement.GetDof();
m_densityShape.SetSize(densityElement.GetDof());
m_displacementShape.SetSize(scalarDisplacementDofCount);
m_referenceDisplacementDShape.SetSize(scalarDisplacementDofCount, dimension);
m_referenceDisplacementJacobian.SetSize(dimension, dimension);
m_physicalPositionVariation.SetSize(dimension);
m_centrifugalAcceleration.SetSize(dimension);
m_centrifugalAccelerationVariation.SetSize(dimension);
m_weightedForce.SetSize(dimension);
m_elementAction.SetSize(data.displacementDofs.Size());
m_elementAction = 0.0;
for (int quadraturePoint = 0; quadraturePoint < data.integrationRule->GetNPoints(); ++quadraturePoint) {
const mfem::IntegrationPoint &integrationPoint = data.integrationRule->IntPoint(quadraturePoint);
densityElement.CalcShape(integrationPoint, m_densityShape);
displacementElement.CalcShape(integrationPoint, m_displacementShape);
displacementElement.CalcDShape(integrationPoint, m_referenceDisplacementDShape);
mfem::MultAtB(directionDofs, m_referenceDisplacementDShape, m_referenceDisplacementJacobian);
directionDofs.MultTranspose(m_displacementShape, m_physicalPositionVariation);
m_rotation->potential_gradient_directional_derivative(
m_physicalPositionVariation, m_centrifugalAccelerationVariation
);
m_centrifugalAccelerationVariation *= -1.0;
double logarithmicJacobianVariation{0.0};
for (int row = 0; row < dimension; ++row) {
m_centrifugalAcceleration(row) = data.centrifugalAccelerations(quadraturePoint, row);
for (int column = 0; column < dimension; ++column) {
logarithmicJacobianVariation +=
data.inverseElementJacobians(quadraturePoint, row * dimension + column) *
m_referenceDisplacementJacobian(column, row);
}
}
const double densityVariationValue = m_elementDensityVariation * m_densityShape;
const double baseDensityValue = data.baseDensityValues(quadraturePoint);
m_weightedForce = 0.0;
m_weightedForce.Add(densityVariationValue, m_centrifugalAcceleration);
m_weightedForce.Add(baseDensityValue, m_centrifugalAccelerationVariation);
m_weightedForce.Add(baseDensityValue * logarithmicJacobianVariation, m_centrifugalAcceleration);
m_weightedForce *= data.quadratureWeights(quadraturePoint);
for (int scalarDof = 0; scalarDof < scalarDisplacementDofCount; ++scalarDof) {
for (int component = 0; component < dimension; ++component) {
const int vectorDof =
vector_dof_index(ordering, scalarDof, component, scalarDisplacementDofCount, dimension);
m_elementAction(vectorDof) += m_displacementShape(scalarDof) * m_weightedForce(component);
}
}
}
if (data.displacementDofTransformation != nullptr) {
data.displacementDofTransformation->TransformDual(m_elementAction);
}
m_localAction.AddElementVector(data.displacementDofs, m_elementAction);
}
local_to_true(*m_fem.displacementFes, m_localAction, actionTrue);
}
void PreparedRotationalDisplacementForceOperator::ApplyCompleteJacobianAction(
const mfem::Vector &densityVariation,
const mfem::Vector &displacementVariation,
@@ -155,10 +416,7 @@ namespace mean_field::operators {
m_context.GetDensityMap().scatter(densityVariation, m_densityVariationTrue);
m_context.GetDisplacementMap().scatter(displacementVariation, m_displacementVariationTrue);
kernels::apply_rotational_displacement_force_complete_action(
m_fem, m_domainMapper, *m_rotation, m_context.GetBaseDensityTrue(), m_densityVariationTrue,
m_displacementVariationTrue, m_context.GetDisplacementTrue(), m_actionTrue
);
ApplyPreparedCompleteJacobianActionTrue(m_densityVariationTrue, m_displacementVariationTrue, m_actionTrue);
action.SetSize(m_context.GetDisplacementMap().reduced_size());
m_context.GetDisplacementMap().gather(m_actionTrue, action);