module; #include #include #include #include #include #include #include #include #include #include module mean_field; import :fem.reference_tables; namespace mean_field::fem { namespace { struct ReferenceTableKey { const mfem::FiniteElement *element; std::vector> points; bool operator<(const ReferenceTableKey &other) const { if (element != other.element) return std::less{}(element, other.element); return points < other.points; } }; ReferenceTableKey make_key( const mfem::FiniteElement &element, const mfem::IntegrationRule &rule ) { ReferenceTableKey key{.element = &element, .points = {}}; key.points.reserve(rule.GetNPoints()); for (int q = 0; q < rule.GetNPoints(); ++q) { const auto &point = rule.IntPoint(q); const std::array values{ point.x, element.GetDim() > 1 ? point.y : 0.0, element.GetDim() > 2 ? point.z : 0.0, point.weight }; for (const double value : values) { if (!std::isfinite(value)) throw std::invalid_argument("Reference table quadrature entries must be finite."); } key.points.push_back(values); } return key; } } // namespace struct ReferenceTableCache::Storage { std::mutex mutex; std::map> scalar_tables; std::map> vector_tables; }; ReferenceTableCache::ReferenceTableCache() : m_storage(std::make_unique()) { } ReferenceTableCache::~ReferenceTableCache() = default; std::shared_ptr ReferenceTableCache::GetScalarTable( const mfem::FiniteElement &element, const mfem::IntegrationRule &rule ) const { if (element.GetRangeType() != mfem::FiniteElement::SCALAR) throw std::invalid_argument("A scalar reference table requires a scalar finite element."); auto key = make_key(element, rule); const std::lock_guard lock(m_storage->mutex); if (const auto found = m_storage->scalar_tables.find(key); found != m_storage->scalar_tables.end()) return found->second; auto table = std::shared_ptr(new ScalarReferenceTable(element, rule)); m_storage->scalar_tables.emplace(std::move(key), table); return table; } std::shared_ptr ReferenceTableCache::GetVectorTable( const mfem::FiniteElement &element, const mfem::IntegrationRule &rule ) const { if (element.GetRangeType() != mfem::FiniteElement::VECTOR) throw std::invalid_argument("A vector reference table requires a vector finite element."); auto key = make_key(element, rule); const std::lock_guard lock(m_storage->mutex); if (const auto found = m_storage->vector_tables.find(key); found != m_storage->vector_tables.end()) return found->second; auto table = std::shared_ptr(new VectorReferenceTable(element, rule)); m_storage->vector_tables.emplace(std::move(key), table); return table; } ScalarReferenceTable::ScalarReferenceTable( const mfem::FiniteElement &element, const mfem::IntegrationRule &rule ) : m_values( rule.GetNPoints(), element.GetDof() ), m_dimension(element.GetDim()) { mfem::Vector values(element.GetDof()); if (element.GetDerivType() == mfem::FiniteElement::GRAD) m_gradients.resize(rule.GetNPoints()); for (int q = 0; q < rule.GetNPoints(); ++q) { const auto &point = rule.IntPoint(q); element.CalcShape(point, values); for (int dof = 0; dof < element.GetDof(); ++dof) m_values(q, dof) = values(dof); if (!m_gradients.empty()) { auto &gradient = m_gradients[q]; gradient.SetSize(element.GetDof(), m_dimension); element.CalcDShape(point, gradient); } } } const mfem::DenseMatrix &ScalarReferenceTable::GetValues() const { return m_values; } const mfem::DenseMatrix &ScalarReferenceTable::GetGradients(const int point) const { return m_gradients.at(point); } int ScalarReferenceTable::GetPointCount() const { return m_values.Height(); } int ScalarReferenceTable::GetDofCount() const { return m_values.Width(); } int ScalarReferenceTable::GetDimension() const { return m_dimension; } VectorReferenceTable::VectorReferenceTable( const mfem::FiniteElement &element, const mfem::IntegrationRule &rule ) : m_dof_count(element.GetDof()), m_dimension(element.GetRangeDim()) { m_values.resize(rule.GetNPoints()); for (int q = 0; q < rule.GetNPoints(); ++q) { auto &values = m_values[q]; values.SetSize(m_dof_count, m_dimension); element.CalcVShape(rule.IntPoint(q), values); } } const mfem::DenseMatrix &VectorReferenceTable::GetValues(const int point) const { return m_values.at(point); } int VectorReferenceTable::GetPointCount() const { return static_cast(m_values.size()); } int VectorReferenceTable::GetDofCount() const { return m_dof_count; } int VectorReferenceTable::GetDimension() const { return m_dimension; } } // namespace mean_field::fem