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:
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user