perf(allocations): reduced overall allocations by 95%, increaseed jacobian applicatin by 2x

This commit uses global pre allocated work space to dramatically reduce memory usage and allocation time
This commit is contained in:
2026-09-10 06:50:56 -04:00
parent b3c04d507a
commit 75cc638739
66 changed files with 207183 additions and 99552 deletions

View File

@@ -311,13 +311,12 @@ namespace mean_field::operators {
const mfem::FiniteElement &densityElement = *m_fem.densityFes->GetFE(elementId);
const mfem::IntegrationRule &integrationRule =
get_moment_of_inertia_rule(m_fem, densityElement, *transformation);
data.quadraturePoints.resize(integrationRule.GetNPoints());
for (int quadraturePoint = 0; quadraturePoint < integrationRule.GetNPoints(); ++quadraturePoint) {
QuadraturePointData &point = data.quadraturePoints[quadraturePoint];
point.integrationPoint = integrationRule.IntPoint(quadraturePoint);
point.densityShape.SetSize(densityElement.GetDof());
densityElement.CalcShape(point.integrationPoint, point.densityShape);
}
data.integrationRule = &integrationRule;
data.densityBasis = m_fem.GetReferenceTables().GetScalarTable(densityElement, integrationRule);
data.mappingContexts.SetSize(integrationRule.GetNPoints(), m_fem.mesh->Dimension());
data.density.SetSize(integrationRule.GetNPoints());
data.quadratureWeights.SetSize(integrationRule.GetNPoints());
data.cylindricalRadiusSquared.SetSize(integrationRule.GetNPoints());
}
int globalStellarElementCount = 0;
MFEM_VERIFY(
@@ -344,6 +343,7 @@ namespace mean_field::operators {
return mapping::MappingStatus::non_finite_result;
}
mapping::DomainMapper::Workspace workspace(m_fem.mesh->Dimension());
mapping::VolumeMappingContext mappingContext;
for (ElementPAData &data : m_elements) {
displacementLocal.GetSubVector(data.displacementDofs, data.baseDisplacement);
@@ -372,21 +372,24 @@ namespace mean_field::operators {
.displacement = displacementData, .compactification = compactificationData
};
mfem::ElementTransformation *transformation = m_fem.mesh->GetElementTransformation(data.elementId);
for (QuadraturePointData &point : data.quadraturePoints) {
for (int quadraturePoint = 0; quadraturePoint < data.integrationRule->GetNPoints(); ++quadraturePoint) {
const mapping::MappingStatus status = m_domainMapper.EvaluateVolume(
mappingData, *transformation, point.integrationPoint, workspace, point.mappingContext
mappingData, *transformation, data.integrationRule->IntPoint(quadraturePoint), workspace,
mappingContext
);
if (status != mapping::MappingStatus::valid) {
return status;
}
if (point.mappingContext.mapping.compactified) {
if (mappingContext.mapping.compactified) {
return mapping::MappingStatus::at_compactified_infinity;
}
point.cylindricalRadiusSquared =
CylindricalRadiusSquared(point.mappingContext.mapping.physical_position);
if (!std::isfinite(point.cylindricalRadiusSquared)) {
data.cylindricalRadiusSquared(quadraturePoint) =
CylindricalRadiusSquared(mappingContext.mapping.physical_position);
if (!std::isfinite(data.cylindricalRadiusSquared(quadraturePoint))) {
return mapping::MappingStatus::non_finite_result;
}
data.mappingContexts.Store(quadraturePoint, mappingContext);
data.quadratureWeights(quadraturePoint) = mappingContext.quadrature.weight;
}
}
return std::nullopt;
@@ -411,11 +414,9 @@ namespace mean_field::operators {
if (!is_finite_vector(elementDensity)) {
return false;
}
for (QuadraturePointData &point : data.quadraturePoints) {
point.density = elementDensity * point.densityShape;
if (!std::isfinite(point.density)) {
return false;
}
data.densityBasis->GetValues().Mult(elementDensity, data.density);
if (!is_finite_vector(data.density)) {
return false;
}
}
return true;
@@ -424,9 +425,9 @@ namespace mean_field::operators {
std::optional<AngularMomentumPreparationRejection> PreparedAngularMomentumOperator::TryAssembleResidual() {
double localMomentOfInertia = 0.0;
for (const ElementPAData &data : m_elements) {
for (const QuadraturePointData &point : data.quadraturePoints) {
localMomentOfInertia +=
point.density * point.cylindricalRadiusSquared * point.mappingContext.quadrature.weight;
for (int quadraturePoint = 0; quadraturePoint < data.density.Size(); ++quadraturePoint) {
localMomentOfInertia += data.density(quadraturePoint) * data.cylindricalRadiusSquared(quadraturePoint) *
data.quadratureWeights(quadraturePoint);
}
}
m_momentOfInertia = GlobalSum(localMomentOfInertia);
@@ -472,15 +473,18 @@ namespace mean_field::operators {
"Angular-momentum density action has the wrong true-vector size."
);
true_to_local(*m_fem.densityFes, densityVariation, m_densityVariationLocal);
mfem::Vector quadratureDensityVariation;
double localAction = 0.0;
for (const ElementPAData &data : m_elements) {
m_densityVariationLocal.GetSubVector(data.densityDofs, m_elementDensityVariation);
if (data.densityDofTransformation != nullptr) {
data.densityDofTransformation->InvTransformPrimal(m_elementDensityVariation);
}
for (const QuadraturePointData &point : data.quadraturePoints) {
localAction += (m_elementDensityVariation * point.densityShape) * point.cylindricalRadiusSquared *
point.mappingContext.quadrature.weight;
quadratureDensityVariation.SetSize(data.integrationRule->GetNPoints());
data.densityBasis->GetValues().Mult(m_elementDensityVariation, quadratureDensityVariation);
for (int quadraturePoint = 0; quadraturePoint < quadratureDensityVariation.Size(); ++quadraturePoint) {
localAction += quadratureDensityVariation(quadraturePoint) *
data.cylindricalRadiusSquared(quadraturePoint) * data.quadratureWeights(quadraturePoint);
}
}
return localAction;
@@ -496,6 +500,7 @@ namespace mean_field::operators {
true_to_local(*m_fem.displacementFes, displacementVariation, m_displacementVariationLocal);
mapping::DomainMapper::Workspace workspace(m_fem.mesh->Dimension());
mapping::VolumeMappingVariation variation;
mapping::VolumeMappingContext mappingContext;
double localAction = 0.0;
for (const ElementPAData &data : m_elements) {
m_displacementVariationLocal.GetSubVector(data.displacementDofs, m_elementDisplacementVariation);
@@ -515,20 +520,22 @@ namespace mean_field::operators {
.displacement = baseDisplacementData, .compactification = compactificationData
};
mfem::ElementTransformation *transformation = m_fem.mesh->GetElementTransformation(data.elementId);
for (const QuadraturePointData &point : data.quadraturePoints) {
for (int quadraturePoint = 0; quadraturePoint < data.integrationRule->GetNPoints(); ++quadraturePoint) {
data.mappingContexts.Load(quadraturePoint, mappingContext);
const mapping::MappingStatus status = m_domainMapper.EvaluateVolumeVariation(
mappingData, directionData, *transformation, point.integrationPoint, point.mappingContext,
workspace, variation
mappingData, directionData, *transformation, data.integrationRule->IntPoint(quadraturePoint),
mappingContext, workspace, variation
);
MFEM_VERIFY(
status == mapping::MappingStatus::valid,
"Mapped angular-momentum variation is invalid. Element: " << data.elementId
);
const double radiusSquaredVariation = CylindricalRadiusSquaredVariation(
point.mappingContext.mapping.physical_position, variation.mapping.physical_position_variation
mappingContext.mapping.physical_position, variation.mapping.physical_position_variation
);
localAction += point.density * (radiusSquaredVariation * point.mappingContext.quadrature.weight +
point.cylindricalRadiusSquared * variation.weight_variation);
localAction += data.density(quadraturePoint) *
(radiusSquaredVariation * data.quadratureWeights(quadraturePoint) +
data.cylindricalRadiusSquared(quadraturePoint) * variation.weight_variation);
}
}
return localAction;