Program Listing for File expr.hpp

Return to documentation for file (SeQuant/core/utility/expr.hpp)

#ifndef SEQUANT_EXPR_UTILITIES_HPP
#define SEQUANT_EXPR_UTILITIES_HPP

#include <SeQuant/core/container.hpp>
#include <SeQuant/core/expr.hpp>
#include <SeQuant/core/index.hpp>
#include <SeQuant/core/utility/expr_matcher.hpp>
#include <SeQuant/core/utility/indices.hpp>
#include <SeQuant/core/utility/macros.hpp>
#include <SeQuant/core/utility/tensor.hpp>

#include <range/v3/algorithm/equal.hpp>
#include <range/v3/algorithm/find.hpp>
#include <range/v3/view/concat.hpp>
#include <range/v3/view/enumerate.hpp>

#include <algorithm>
#include <cassert>
#include <concepts>
#include <optional>
#include <string>
#include <string_view>
#include <type_traits>
#include <utility>

namespace sequant {

std::string diff(const Expr &lhs, const Expr &rhs);

bool is_valid(const ExprPtr &expr, std::string *msg = nullptr);

bool is_valid(const Expr &expr, std::string *msg = nullptr);

bool is_valid(const ResultExpr &expr, std::string *msg = nullptr);

[[nodiscard]] ExprPtr transform_expr(
    const ExprPtr &expr, const container::map<Index, Index> &index_replacements,
    Constant::scalar_type scaling_factor = 1);
[[nodiscard]] ExprPtr transform_expr(
    const Expr &expr, const container::map<Index, Index> &index_replacements,
    Constant::scalar_type scaling_factor = 1);

std::optional<ExprPtr> pop_tensor(ExprPtr &expression, std::wstring_view label);

ExprPtr &replace(ExprPtr &expr, const ExprMatcher &target,
                 const Expr &replacement);

template <typename EqualityComparator = std::equal_to<>>
[[deprecated(
    "This is only a backwards-compat shim. Use overload using "
    "ExprMatcher instead")]] ExprPtr &
replace(ExprPtr &expr, const ExprPtr &target, const ExprPtr &replacement,
        EqualityComparator = {}) {
  ExprMatcherOptions options{.cross_comparisons = true};
  if constexpr (std::same_as<std::remove_cvref_t<EqualityComparator>,
                             std::equal_to<>>) {
    options.tensor_cmp = TensorComparison::Identity;
  } else if constexpr (std::same_as<std::remove_cvref_t<EqualityComparator>,
                                    TensorBlockEqualComparator>) {
    options.tensor_cmp = TensorComparison::Block;
  } else {
    static_assert(false,
                  "Compatibility shim can't deal with the provided comparator");
  }

  return replace(expr, ExprMatcher(*target, std::move(options)), *replacement);
}

ResultExpr &replace(ResultExpr &expr, const ExprMatcher &target,
                    const Expr &replacement);

template <typename EqualityComparator = std::equal_to<>>
[[deprecated(
    "This is only a backwards-compat shim. Use overload using "
    "ExprMatcher instead")]] ResultExpr &
replace(ResultExpr &expr, const ExprPtr &target, const ExprPtr &replacement,
        EqualityComparator = {}) {
  ExprMatcherOptions options{.cross_comparisons = true};
  if constexpr (std::same_as<std::remove_cvref_t<EqualityComparator>,
                             std::equal_to<>>) {
    options.tensor_cmp = TensorComparison::Identity;
  } else if constexpr (std::same_as<std::remove_cvref_t<EqualityComparator>,
                                    TensorBlockEqualComparator>) {
    options.tensor_cmp = TensorComparison::Block;
  } else {
    static_assert(false,
                  "Compatibility shim can't deal with the provided comparator");
  }

  return replace(expr, ExprMatcher(*target, std::move(options)), *replacement);
}

}  // namespace sequant

#endif