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:
@@ -420,17 +420,12 @@ namespace mean_field::operators {
|
||||
|
||||
const mfem::IntegrationRule &integrationRule =
|
||||
get_mass_normalization_rule(m_fem, densityElement, *transformation);
|
||||
data.integrationRule = &integrationRule;
|
||||
|
||||
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.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());
|
||||
}
|
||||
|
||||
int globalStellarElementCount = 0;
|
||||
@@ -459,6 +454,7 @@ namespace mean_field::operators {
|
||||
true_to_local(*m_fem.displacementFes, displacement, displacementLocal);
|
||||
|
||||
mapping::DomainMapper::Workspace workspace(m_fem.mesh->Dimension());
|
||||
mapping::VolumeMappingContext mappingContext;
|
||||
std::optional<MassNormalizationPreparationRejection> rejection;
|
||||
|
||||
for (ElementPAData &data : m_elements) {
|
||||
@@ -491,9 +487,10 @@ namespace mean_field::operators {
|
||||
|
||||
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
|
||||
);
|
||||
|
||||
MFEM_VERIFY(
|
||||
@@ -505,6 +502,9 @@ namespace mean_field::operators {
|
||||
rejection, {.reason = MassNormalizationPreparationRejectionReason::mapping_failure,
|
||||
.mappingStatus = status}
|
||||
);
|
||||
} else {
|
||||
data.mappingContexts.Store(quadraturePoint, mappingContext);
|
||||
data.quadratureWeights(quadraturePoint) = mappingContext.quadrature.weight;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -537,9 +537,9 @@ namespace mean_field::operators {
|
||||
data.densityDofTransformation->InvTransformPrimal(elementDensity);
|
||||
}
|
||||
|
||||
for (QuadraturePointData &point : data.quadraturePoints) {
|
||||
point.density = elementDensity * point.densityShape;
|
||||
if (!std::isfinite(point.density)) {
|
||||
data.densityBasis->GetValues().Mult(elementDensity, data.density);
|
||||
for (int quadraturePoint = 0; quadraturePoint < data.density.Size(); ++quadraturePoint) {
|
||||
if (!std::isfinite(data.density(quadraturePoint))) {
|
||||
retain_higher_priority_rejection(
|
||||
rejection,
|
||||
{.reason = MassNormalizationPreparationRejectionReason::non_finite_density_interpolation}
|
||||
@@ -556,8 +556,8 @@ namespace mean_field::operators {
|
||||
std::optional<MassNormalizationPreparationRejection> localRejection;
|
||||
|
||||
for (const ElementPAData &data : m_elements) {
|
||||
for (const QuadraturePointData &point : data.quadraturePoints) {
|
||||
const double contribution = point.density * point.mappingContext.quadrature.weight;
|
||||
for (int quadraturePoint = 0; quadraturePoint < data.density.Size(); ++quadraturePoint) {
|
||||
const double contribution = data.density(quadraturePoint) * data.quadratureWeights(quadraturePoint);
|
||||
if (!std::isfinite(contribution) || !std::isfinite(localMass + contribution)) {
|
||||
localRejection = {.reason = MassNormalizationPreparationRejectionReason::non_finite_assembled_mass};
|
||||
continue;
|
||||
@@ -615,6 +615,7 @@ namespace mean_field::operators {
|
||||
true_to_local(*m_fem.densityFes, densityVariation, densityVariationLocal);
|
||||
|
||||
mfem::Vector elementDensityVariation;
|
||||
mfem::Vector quadratureDensityVariation;
|
||||
double localAction = 0.0;
|
||||
|
||||
for (const ElementPAData &data : m_elements) {
|
||||
@@ -624,8 +625,10 @@ namespace mean_field::operators {
|
||||
data.densityDofTransformation->InvTransformPrimal(elementDensityVariation);
|
||||
}
|
||||
|
||||
for (const QuadraturePointData &point : data.quadraturePoints) {
|
||||
localAction += (elementDensityVariation * point.densityShape) * point.mappingContext.quadrature.weight;
|
||||
quadratureDensityVariation.SetSize(data.integrationRule->GetNPoints());
|
||||
data.densityBasis->GetValues().Mult(elementDensityVariation, quadratureDensityVariation);
|
||||
for (int quadraturePoint = 0; quadraturePoint < quadratureDensityVariation.Size(); ++quadraturePoint) {
|
||||
localAction += quadratureDensityVariation(quadraturePoint) * data.quadratureWeights(quadraturePoint);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -650,6 +653,7 @@ namespace mean_field::operators {
|
||||
|
||||
mapping::DomainMapper::Workspace workspace(m_fem.mesh->Dimension());
|
||||
mapping::VolumeMappingVariation variation;
|
||||
mapping::VolumeMappingContext mappingContext;
|
||||
|
||||
mfem::Vector elementDisplacementVariation;
|
||||
double localAction = 0.0;
|
||||
@@ -681,10 +685,11 @@ namespace mean_field::operators {
|
||||
|
||||
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(
|
||||
@@ -694,7 +699,7 @@ namespace mean_field::operators {
|
||||
<< ", status: " << static_cast<int>(status)
|
||||
);
|
||||
|
||||
localAction += point.density * variation.weight_variation;
|
||||
localAction += data.density(quadraturePoint) * variation.weight_variation;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -789,14 +794,13 @@ namespace mean_field::operators {
|
||||
localDual = 0.0;
|
||||
|
||||
mfem::Vector elementDual;
|
||||
mfem::Vector weightedDual;
|
||||
|
||||
for (const ElementPAData &data : m_elements) {
|
||||
elementDual.SetSize(data.densityDofs.Size());
|
||||
elementDual = 0.0;
|
||||
|
||||
for (const QuadraturePointData &point : data.quadraturePoints) {
|
||||
elementDual.Add(residualDual * point.mappingContext.quadrature.weight, point.densityShape);
|
||||
}
|
||||
weightedDual = data.quadratureWeights;
|
||||
weightedDual *= residualDual;
|
||||
data.densityBasis->GetValues().MultTranspose(weightedDual, elementDual);
|
||||
|
||||
if (data.densityDofTransformation != nullptr) {
|
||||
data.densityDofTransformation->TransformDual(elementDual);
|
||||
@@ -821,6 +825,7 @@ namespace mean_field::operators {
|
||||
|
||||
mapping::DomainMapper::Workspace workspace(m_fem.mesh->Dimension());
|
||||
mapping::VolumeMappingVariation variation;
|
||||
mapping::VolumeMappingContext mappingContext;
|
||||
mfem::Vector elementDirection;
|
||||
mfem::Vector elementDual;
|
||||
|
||||
@@ -852,10 +857,11 @@ namespace mean_field::operators {
|
||||
|
||||
double elementDofAction = 0.0;
|
||||
|
||||
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(
|
||||
@@ -864,7 +870,7 @@ namespace mean_field::operators {
|
||||
<< data.elementId << ", status: " << static_cast<int>(status)
|
||||
);
|
||||
|
||||
elementDofAction += point.density * variation.weight_variation;
|
||||
elementDofAction += data.density(quadraturePoint) * variation.weight_variation;
|
||||
}
|
||||
|
||||
elementDual(elementDof) = residualDual * elementDofAction;
|
||||
|
||||
Reference in New Issue
Block a user