Program Listing for File legality.hpp

Return to documentation for file (SeQuant/core/eval/legality.hpp)

#ifndef SEQUANT_EVAL_LEGALITY_HPP
#define SEQUANT_EVAL_LEGALITY_HPP

#include <SeQuant/core/batch_policy.hpp>
#include <SeQuant/core/container.hpp>
#include <SeQuant/core/eval/eval.hpp>
#include <SeQuant/core/eval/eval_expr.hpp>
#include <SeQuant/core/eval/peak_profile.hpp>
#include <SeQuant/core/index.hpp>
#include <SeQuant/core/utility/macros.hpp>

#include <algorithm>
#include <cstddef>
#include <cstdlib>
#include <functional>
#include <iostream>
#include <ranges>
#include <string>
#include <type_traits>
#include <unordered_map>
#include <utility>
#include <vector>

namespace sequant::eval {


enum class LoopRole {
  LoopLocal,
  Reduction,
  LoopCarried,
  LoopInvariant,
};

struct AxisClass {
  Index axis;
  LoopRole role;
};

struct CellLegality {
  std::size_t hash = 0;

  container::svector<Index> build_site;

  container::svector<AxisClass> per_axis;

  container::svector<Index> home_floor;

  container::svector<Index> forced_split_axes;
};

struct LegalitySchedule {
  container::svector<CellLegality> cells;
};

[[nodiscard]] inline container::svector<Index> forced_split_types(
    CellLegality const& cell) {
  container::svector<Index> out;
  for (Index const& ix : cell.forced_split_axes) {
    auto const same_type = [&](Index const& o) {
      return o.space().base_key() == ix.space().base_key();
    };
    if (std::find_if(out.begin(), out.end(), same_type) == out.end())
      out.push_back(ix);
  }
  return out;
}

[[nodiscard]] inline container::svector<Index> build_site_of(
    meta::eval_node auto const& node, BatchPolicy const& policy) {
  container::svector<Index> result;
  auto const add_if_new = [&](Index const& ix) {
    if (std::find(result.begin(), result.end(), ix) == result.end())
      result.push_back(ix);
  };

  // Source the build-site axes from the cost model's actual per-node decision,
  // not from policy.is_batchable_index() (what could be batched). Using the
  // predicate made the ordered schedule loop over every batchable axis (e.g.
  // aux whenever batch:aux_target_size>0) at every node carrying/contracting
  // it, regardless of whether the peak-constrained optimizer decided to slice
  // it -- always-on batching independent of peak_threshold, and over-scoping
  // relative to the DP even under a finite budget. The two authoritative
  // sources are:
  //   - a result (carried) index is a build-site axis iff it slices this node's
  //     own result slots per the cross-occurrence meet (\c sliced_modes, which
  //     already folds in ancestor batch loops the value is variant to); and
  //   - a contracted-at-node index iff the DP batched it here (a \c Contracted
  //     \c node_slice_mask stamp == the node's chosen aprime).
  // With an infinite budget (or no batchable axis) the DP emits no stamps, so
  // both are empty and the schedule is flat -- matching forest descent.
  auto const& sliced = sequant::home_scope(node);
  auto const& stamps = node->node_slice_mask();
  for (Index const& ix : node->canon_indices())
    if (std::find(sliced.begin(), sliced.end(), ix) != sliced.end())
      add_if_new(ix);
  for (Index const& ix : contracted_indices(node))
    if (std::any_of(stamps.begin(), stamps.end(), [&](auto const& p) {
          return p.second == BatchModeType::Contracted && p.first == ix;
        }))
      add_if_new(ix);
  (void)policy;
  return result;
}

[[nodiscard]] inline LoopRole classify_axis(
    container::svector<Index> const& carried,
    container::svector<Index> const& contracted_below, Index const& axis,
    container::svector<OccurrenceRec> const& occurrences,
    container::svector<Index> const& sliced,
    std::function<int(OccurrenceRec const&, Index const&)> const&
        enclosing_slot = {}) {
  auto const same_type = [&](Index const& ix) {
    return ix.space().base_key() == axis.space().base_key();
  };
  // A carried same-space index is a batched loop mode (subject to the lockstep
  // test below) iff it is one of the value's sliced modes; otherwise it is a
  // free full "spectator" dimension (e.g. a retained occ index the DP did not
  // batch) that has no loop at all and must not make the axis LoopCarried --
  // the value is still loop-local w.r.t. the axis's own loop, carrying the
  // spectator dimension through as full. Compared by identity (same
  // ordinal/proto), so a spectator i_4 beside a batched i_1 is skipped.
  auto const is_batched = [&](Index const& ix) {
    return std::any_of(sliced.begin(), sliced.end(),
                       [&](Index const& s) { return s == ix; });
  };

  // The role is decided for this axis, by identity, not for its space: a
  // value may carry one index of a space (loop-local on that instance) and
  // contract another index of the same space in batches (a reduction over a
  // different instance) -- e.g. a residual-pair intermediate carrying the
  // external pair ij while reducing over a contracted pair kl, once the
  // occupied space is batchable in both roles. Deciding by space conflated
  // the two: the carried branch won and the reduction was never recorded, so
  // the builder emitted a scatter escape at the contracted instance with no
  // sliced position to scatter ("scatters nothing: empty scatter map").
  auto const carried_pos = std::find(carried.begin(), carried.end(), axis);
  if (carried_pos == carried.end()) {
    bool const reduces_axis =
        std::find(contracted_below.begin(), contracted_below.end(), axis) !=
        contracted_below.end();
    return reduces_axis ? LoopRole::Reduction : LoopRole::LoopInvariant;
  }
  // The carried slot this axis occupies; every occurrence's carried list is
  // positionally aligned with the value's canonical carried list (each
  // occurrence's frame relabels the same slots), so the slot, not the label,
  // identifies the axis across occurrences.
  std::size_t const axis_slot =
      static_cast<std::size_t>(carried_pos - carried.begin());

  bool found_enclosing = false;
  for (OccurrenceRec const& occ : occurrences) {
    // All same-type enclosing loops at this occurrence, not just the first:
    // nested same-space loops each bind a distinct carried index of that
    // type (i_1 under the outer loop, i_2 under the inner). Matching every
    // carried slot against only the first enclosing loop mis-flags a value
    // that carries two same-space indices, each lockstep with its own nested
    // loop, as LoopCarried -- because the second carried index (i_2) never
    // equals the first loop's Index (i_1). Collect the whole same-type
    // enclosing set and let each carried slot lock to any member.
    container::svector<Index> encl;
    for (auto const& e : occ.ectx)
      if (same_type(e.first)) encl.push_back(e.first);
    if (encl.empty())
      continue;  // no enclosing loop of this type at this occurrence
    found_enclosing = true;

    // A carried slot is lockstep iff it equals some enclosing loop's Index
    // (any nesting level). Only an occurrence whose every same-type carried
    // slot is lockstep with some enclosing loop is truly loop-local; a single
    // carried slot with no matching enclosing loop (a free / cross-iteration
    // read) makes the whole occurrence -- and thus the axis -- LoopCarried.
    bool matched_any_same_type = false;
    // Loop-instance test at the set level: the fusion slots of the value's
    // own same-space batched positions at this occurrence. An enclosing loop
    // whose instance is not among them is a loop this value is not sliced by
    // in any position -- a sibling nest (e.g. a 4-occupied intermediate
    // built in nest {s3,s4} and read from nest {s5,s6}) -- so the value must
    // be an assembled escape of its own nest: LoopCarried. A permutation of
    // the same instances (a symmetric value read with i_1/i_2 transposed,
    // slots {1,2} either way) is not this case: the value-keyed cache serves
    // it as a distinct colored cell in-nest, and forcing it to the root would
    // strand the in-nest read.
    // The value's production instances: the front occurrence's stamps (the
    // value is built there; a read occurrence's own stamps follow the
    // consumer's instance at that read -- a sibling-nest read is stamped with
    // the sibling's slots and would pass a self-comparison).
    container::svector<int> own_slots;
    {
      OccurrenceRec const& prod = occurrences.front();
      std::size_t const pc = axis_slot;
      if (pc < prod.carried.size() && same_type(prod.carried[pc]) &&
          is_batched(prod.carried[pc]) && pc < prod.loop_slot.size() &&
          prod.loop_slot[pc] >= 0)
        own_slots.push_back(prod.loop_slot[pc]);
    }
    // Direction matters: the value may be read inside loops it is invariant
    // to (a deeper level of its own nest -- enclosing instances beyond its
    // own are fine), but every one of its own instances must be among the
    // enclosing loops, else it is being read outside its own loops (a sibling
    // nest) and must be an assembled escape. Skipped when any enclosing
    // instance cannot be resolved (incomplete evidence -> label-only test).
    if (enclosing_slot && !own_slots.empty()) {
      container::svector<int> encl_slots;
      bool complete = true;
      for (Index const& L : encl) {
        int const es = enclosing_slot(occ, L);
        if (es < 0) {
          complete = false;
          break;
        }
        encl_slots.push_back(es);
      }
      if (complete)
        for (int os : own_slots)
          if (std::find(encl_slots.begin(), encl_slots.end(), os) ==
              encl_slots.end())
            return LoopRole::LoopCarried;  // own loop instance not enclosing
    }
    if (axis_slot < occ.carried.size()) {
      Index const& c = occ.carried[axis_slot];
      // Batched-ness of this occurrence's position is read off its own home
      // (same frame as `c`), never off `sliced`, which is labeled in the
      // representative occurrence's frame: one node's occurrences carry
      // different labels at one position across terms (i_3 here, i_4
      // there), and a label miss here wrongly made a value every consumer
      // reads inside its loop LoopCarried -- which then bumped its
      // consumers into later passes and materialized it whole.
      bool const batched_here =
          std::find(occ.home.begin(), occ.home.end(), c) != occ.home.end();
      if (same_type(c) && batched_here) {
        matched_any_same_type = true;
        bool const lockstep = std::any_of(
            encl.begin(), encl.end(), [&](Index const& L) { return c == L; });
        if (!lockstep)
          return LoopRole::LoopCarried;  // free / cross-iteration binding
      }
    }
    if (!matched_any_same_type)
      return LoopRole::LoopCarried;  // this slot is not a batched loop slot
  }
  return found_enclosing ? LoopRole::LoopLocal : LoopRole::LoopCarried;
}

template <meta::eval_node_range R>
[[nodiscard]] inline LegalitySchedule analyze_legality(
    RichSchedule const& rich, R const& forest, BatchPolicy const& policy) {
  using Node = std::ranges::range_value_t<R>;

  // point -> occurrence, to reach the parent occurrence (same tree) of an
  // occurrence and read the fusion slot of an enclosing loop by its
  // tree-frame label (carried position, or a reduced mode of the parent).
  std::unordered_map<std::size_t, OccurrenceRec const*> point_occ;
  for (ValueCell const& vc : rich.cells)
    for (OccurrenceRec const& occ : vc.occurrences) point_occ[occ.point] = &occ;
  std::function<int(OccurrenceRec const&, Index const&)> const enclosing_slot =
      [&point_occ](OccurrenceRec const& occ, Index const& L) -> int {
    // Walk up the tree: the loop labeled L is opened by some ancestor; the
    // nearest ancestor carrying (or reducing) L in its own frame stamps its
    // slot.
    std::size_t pt = occ.consumer_point;
    for (int guard = 0; guard < 64; ++guard) {
      auto const it = point_occ.find(pt);
      if (it == point_occ.end()) return -1;
      OccurrenceRec const& par = *it->second;
      for (std::size_t p = 0; p < par.carried.size(); ++p)
        if (par.carried[p] == L)
          return p < par.loop_slot.size() ? par.loop_slot[p] : -1;
      for (auto const& [rm, rs] : par.reduced_slot)
        if (rm == L) return rs;
      if (par.consumer_point == par.point) return -1;  // forest root
      pt = par.consumer_point;
    }
    return -1;
  };

  // One representative forest node per value (keyed by value id, not node
  // id): a value's roles are read off one of its own occurrences' nodes -- a
  // node of the same hash home-sliced elsewhere carries that frame's slice
  // annotations, not this value's.
  //
  // Held by pointer, not by value: \c Node's copy constructor deep-copies the
  // whole subtree, so one entry per node costs O(subtree) and a by-value map is
  // O(n^2) in forest size -- gigabytes on a thousand-summand spine, which is
  // OOM long before any stack limit matters. \p forest outlives
  // this call (it is the caller's), so the pointees stay valid; the entries are
  // read-only here.
  //
  // Taking addresses requires the range to yield references to the caller's
  // nodes; a range of prvalues (a transform view, say) would hand us dangling
  // pointers. Same guard, same reason, as cache_manager()'s pointer-keyed DAG
  // walk (cache_manager.hpp) and the value->node bridges
  // (value_node_map.hpp).
  // `R const`, not `R`: the walk below iterates `forest`, which is bound as
  // `R const&`, so the const-qualified range's reference type is the one that
  // has to be a reference. A range that yields references when mutable and
  // prvalues when const would slip past the unqualified form.
  static_assert(
      std::is_reference_v<std::ranges::range_reference_t<R const>>,
      "analyze_legality(): the forest range must yield references to nodes "
      "that outlive the call (the value->node map keys on their addresses)");
  std::unordered_map<std::size_t, Node const*> node_of;
  {
    // Iterative pre-order: the residual's Sum spine is as deep as the number
    // of terms, so a recursive descent would overflow the call stack. Right
    // child pushed before left, so the pop order is the recursion's.
    std::vector<Node const*> stack;
    for (auto const& tree : forest) {
      stack.push_back(&tree);
      while (!stack.empty()) {
        Node const& n = *stack.back();
        stack.pop_back();
        node_of.emplace(value_key_of(n), &n);
        if (!n.leaf()) {
          stack.push_back(&n.right());
          stack.push_back(&n.left());
        }
      }
    }
  }

  (void)policy;  // build-site now sourced from the DP decision, not the policy

  // The single classification round over every cell.
  auto const build_cells = [&]() -> LegalitySchedule {
    LegalitySchedule out;
    out.cells.reserve(rich.cells.size());
    for (ValueCell const& vc : rich.cells) {
      CellLegality cl;
      cl.hash = value_key_of(vc);

      // Every RichSchedule cell was produced by walking this same forest (see
      // compute_dag_boulevard), so its hash must resolve here; a miss would
      // silently leave contracted_below empty and misclassify the cell rather
      // than surface the underlying bug.
      auto const it = node_of.find(value_key_of(vc));
      SEQUANT_ASSERT(it != node_of.end());
      container::svector<Index> contracted_below;
      for (Index const& ix : contracted_indices(*it->second))
        contracted_below.push_back(ix);

      // Build-site axes come from the cost model's actual per-node decision,
      // not policy.is_batchable_index() (what could be batched). Sourcing from
      // the predicate made the ordered schedule loop over every batchable axis
      // (e.g. aux whenever batch:aux_target_size>0) regardless of whether the
      // peak-constrained optimizer sliced it -- always-on batching independent
      // of peak_threshold, and over-scoping relative to the DP under a finite
      // budget. A result (carried) axis is a build-site axis iff it slices this
      // node's own result slots per the cross-occurrence meet (sliced_modes,
      // which folds in the ancestor batch loops the value is variant to); a
      // contracted-at-node axis iff the DP batched it here (a Contracted
      // node_slice_mask stamp == the node's chosen aprime). With an infinite
      // budget (or no batchable axis) the DP emits no stamps, so the site is
      // empty and the schedule is flat -- identical to forest descent.
      container::svector<Index> site;
      auto const add_if_new = [&](Index const& ix) {
        if (std::find(site.begin(), site.end(), ix) == site.end())
          site.push_back(ix);
      };
      auto const& dp_sliced = sequant::home_scope(*it->second);
      auto const& dp_stamps = (*it->second)->node_slice_mask();
      // Positional: the representative node's home is labeled in its tree's
      // frame, vc.carried in the first occurrence's; positions are canonical
      // across occurrences (explicit-cells design section 11), labels are
      // not.
      {
        auto const& rep_carried = (*it->second)->canon_indices();
        for (std::size_t p = 0; p < vc.carried.size() && p < rep_carried.size();
             ++p)
          if (std::find(dp_sliced.begin(), dp_sliced.end(), rep_carried[p]) !=
              dp_sliced.end())
            add_if_new(vc.carried[p]);
      }
      for (Index const& ix : contracted_below)
        if (std::any_of(dp_stamps.begin(), dp_stamps.end(), [&](auto const& p) {
              return p.second == BatchModeType::Contracted && p.first == ix;
            }))
          add_if_new(ix);

      for (Index const& axis : site) {
        AxisClass ac;
        ac.axis = axis;
        ac.role = classify_axis(vc.carried, contracted_below, axis,
                                vc.occurrences, dp_sliced, enclosing_slot);
        cl.per_axis.push_back(std::move(ac));
      }

      for (AxisClass const& ac : cl.per_axis) {
        // home_floor: the LoopLocal subset of per_axis -- the axes the value
        // stays sliced on (homed inside). Every other build-site axis
        // (Reduction, LoopCarried) is lifted out, as is any axis outside
        // build_site altogether (the implicit LoopInvariant case).
        if (ac.role == LoopRole::LoopLocal) cl.home_floor.push_back(ac.axis);
        // forced_split_axes: the LoopCarried subset. The value survives the
        // axis into its own result, so its producing loop must close before
        // any cross-iteration consumer -- the loop cannot stay a single
        // unbroken pass around this value.
        if (ac.role == LoopRole::LoopCarried)
          cl.forced_split_axes.push_back(ac.axis);
      }

      cl.build_site = std::move(site);
      out.cells.push_back(std::move(cl));
    }
    return out;
  };

  return build_cells();
}

}  // namespace sequant::eval

#endif  // SEQUANT_EVAL_LEGALITY_HPP