Program Listing for File MoveBucket.h

Return to documentation for file (src/partition/refine/bucket/MoveBucket.h)

#ifndef PARTITION_REFINE_BUCKET_MOVEBUCKET_H_
#define PARTITION_REFINE_BUCKET_MOVEBUCKET_H_

#include "refine/core/RefineTypes.h"


// ---------------------------------------------------------------------------
// MovePriorityQueue<KeyType> — unified binary-heap priority queue.
//
//   KeyType = int              → per-part bucket (former MyPriorityQueue)
//             pair<int,int>    → global single bucket (former
//             MyPriorityQueueAllInOne)
//
// The map backend is selected at compile time:
//   int key  → vector<int>  (direct index, sentinel -1)
//   pair key → unordered_map<pair<int,int>, int, hashfunc>
// ---------------------------------------------------------------------------

template <typename KeyType>
class MovePriorityQueue {
  // ---- map type selection ------------------------------------------------
  using MapType = conditional_t<is_same_v<KeyType, int>, vector<int>,
                                unordered_map<pair<int, int>, int, hashfunc>>;

  // ---- key extraction from a Gain element --------------------------------
  static KeyType keyOf(const Gain &g)
  {
    if constexpr (is_same_v<KeyType, int>) {
      return g.get_id();
    } else {
      return make_pair(g.get_destination(), g.get_id());
    }
  }

  // ---- map helpers -------------------------------------------------------
  int mapGet(const KeyType &k) const
  {
    if constexpr (is_same_v<KeyType, int>) {
      return vertices_map_[k];
    } else {
      auto it = vertices_map_.find(k);
      return (it != vertices_map_.end()) ? it->second : -1;
    }
  }
  void mapSet(const KeyType &k, int v)
  {
    vertices_map_[k] = v;
  }
  void mapErase(const KeyType &k)
  {
    if constexpr (is_same_v<KeyType, int>) {
      vertices_map_[k] = -1;
    } else {
      vertices_map_.erase(k);
    }
  }
  bool mapHas(const KeyType &k) const
  {
    if constexpr (is_same_v<KeyType, int>) {
      return k >= 0 && k < static_cast<int>(vertices_map_.size()) &&
             vertices_map_[k] != -1;
    } else {
      return vertices_map_.find(k) != vertices_map_.end();
    }
  }
  void mapClear()
  {
    if constexpr (is_same_v<KeyType, int>) {
      fill(vertices_map_.begin(), vertices_map_.end(), -1);
    } else {
      vertices_map_.clear();
    }
  }

 public:
  MovePriorityQueue(const int total, const int max_level, const graph &g)
      : max_level_(max_level), g_(g)
  {
    if constexpr (is_same_v<KeyType, int>) {
      vertices_map_.resize(total);
      fill(vertices_map_.begin(), vertices_map_.end(), -1);
    }
    total_  = 0;
    active_ = false;
  }

  void clear()
  {
    mapClear();
    total_ = 0;
    gains_.clear();
    active_ = false;
  }

  bool compare_element(int id_a, int id_b) const
  {
    if (gains_[id_a].get_gain() != gains_[id_b].get_gain()) {
      return gains_[id_a].get_gain() > gains_[id_b].get_gain();
    }
    return g_.nodes[gains_[id_a].get_id()].weight <
           g_.nodes[gains_[id_b].get_id()].weight;
  }

  void heap_up(int id)
  {
    while (id > 0 && compare_element(id, Parent(id))) {
      auto &parent = gains_[Parent(id)];
      auto &child  = gains_[id];
      mapSet(keyOf(child), Parent(id));
      mapSet(keyOf(parent), id);
      swap(parent, child);
      id = Parent(id);
    }
  }

  void heap_down(int id)
  {
    int max_id        = id;
    const int left_id = Left(id);
    if (left_id < total_ && compare_element(left_id, max_id)) {
      max_id = left_id;
    }
    const int right_id = Right(id);
    if (right_id < total_ && compare_element(right_id, max_id)) {
      max_id = right_id;
    }
    if (id == max_id) {
      return;
    }
    auto &cur = gains_[id];
    auto &max = gains_[max_id];
    mapSet(keyOf(cur), max_id);
    mapSet(keyOf(max), id);
    swap(cur, max);
    heap_down(max_id);
  }

  void change_priority(KeyType id, const Gain &new_gain)
  {
    const int index = mapGet(id);
    if (index == -1) {
      return;
    }
    const double old_gain = gains_[index].get_gain();
    gains_[index]         = new_gain;
    if (new_gain.get_gain() > old_gain) {
      heap_up(index);
    } else {
      heap_down(index);
    }
  }

  void insert(const Gain &ele)
  {
    if (total_ == 0) {
      active_ = true;
    }
    total_++;
    gains_.push_back(ele);
    mapSet(keyOf(ele), total_ - 1);
    heap_up(total_ - 1);
  }

  Gain top() const
  {
    return gains_.front();
  }

  Gain pop()
  {
    auto max_ele = gains_.front();
    gains_[0]    = gains_[total_ - 1];
    mapSet(keyOf(gains_[total_ - 1]), 0);
    total_--;
    gains_.pop_back();
    heap_down(0);
    mapErase(keyOf(max_ele));
    return max_ele;
  }

  void remove(KeyType id)
  {
    const int index = mapGet(id);
    if (index == -1) {
      return;
    }
    gains_[index].set_gain(gains_.front().get_gain() + 1.0);
    heap_up(index);
    pop();
    if (total_ <= 0) {
      active_ = false;
    }
  }

  Gain get(const vector<VectorXi> &occupied_resources, const graph &g,
           bool bound_constraint, vector<VectorXi> fpga_resources,
           VectorXi upper_resources, VectorXi lower_resources)
  {
    if (total_ <= 0) {
      return Gain();
    }
    int pass      = 0;
    int candidate = -1;
    int index     = 0;
    auto check    = [&](int index_id) {
      const int id  = gains_[index_id].get_id();
      const int dst = gains_[index_id].get_destination();
      const VectorXi &actual_upper_resources =
          bound_constraint ? upper_resources : fpga_resources[dst];
      int gap1 = (actual_upper_resources - occupied_resources[dst] -
                  g.nodes[id].resources)
                     .minCoeff();
      if (bound_constraint) {
        const int src = gains_[index_id].get_source();
        int gap2 =
            (occupied_resources[src] - g.nodes[id].resources - lower_resources)
                .minCoeff();
        return gap1 >= 0 && gap2 >= 0;
      } else {
        return gap1 >= 0;
      }
    };
    if (check(index)) {
      return gains_[index];
    }
    while (pass < max_level_) {
      pass++;
      const int left = Left(index);
      if (left < total_ && check(left)) {
        candidate = left;
      }
      const int right = Right(index);
      if (right < total_ && check(right) &&
          (candidate == -1 || compare_element(right, candidate))) {
        candidate = right;
      }
      if (candidate > 0) {
        return gains_[candidate];
      }
      if (left >= total_ || right >= total_) {
        return Gain();
      }
      index = compare_element(right, left) ? right : left;
    }
    return Gain();
  }

  bool check(KeyType id) const
  {
    return mapHas(id);
  }

  bool get_active() const
  {
    return active_;
  }
  void set_active(const bool active)
  {
    active_ = active;
  }
  int get_total() const
  {
    return total_;
  }

 private:
  bool active_;
  const graph &g_;
  int total_     = 0;
  int max_level_ = 25;
  MapType vertices_map_;
  vector<Gain> gains_;

  static int Parent(int element)
  {
    return (element - 1) / 2;
  }
  static int Left(int element)
  {
    return 2 * element + 1;
  }
  static int Right(int element)
  {
    return 2 * element + 2;
  }
};

// Type aliases preserving the original names used throughout the codebase.
using MyPriorityQueue =
    MovePriorityQueue<int>;
using MyPriorityQueueAllInOne = MovePriorityQueue<pair<
    int,
    int>>;

// ---------------------------------------------------------------------------
// Bucket wrappers used by FM/DSFM strategy pipelines.
// ---------------------------------------------------------------------------

class PerPartBucketSet {
 public:
  PerPartBucketSet(int part_count, int node_count, int max_level,
                   const graph &g);

  void insert(int part_id, const Gain &gain);
  bool contains(int part_id, int vertex_id) const;
  void changePriority(int part_id, int vertex_id, const Gain &gain);
  void upsert(int part_id, int vertex_id, const Gain &gain);
  Gain pop(int part_id);
  Gain top(int part_id) const;
  Gain get(int part_id, const vector<VectorXi> &occupied_resources,
           const graph &g, bool bound_constraint,
           const vector<VectorXi> &fpga_resources, const VectorXi &upper,
           const VectorXi &lower);
  void remove(int part_id, int vertex_id);
  void removeFromAll(int vertex_id);
  bool isActive(int part_id) const;
  int total(int part_id) const;
  int size() const;
  void clear();
  void clear(int part_id);

 private:
  vector<MyPriorityQueue> buckets_;
};

class GlobalMoveBucket {
 public:
  GlobalMoveBucket(int total_size, int max_level, const graph &g);

  void insert(const Gain &gain);
  bool contains(const pair<int, int> &key) const;
  bool contains(int part_id, int vertex_id) const;
  void changePriority(const pair<int, int> &key, const Gain &gain);
  void changePriority(int part_id, int vertex_id, const Gain &gain);
  void upsert(int part_id, int vertex_id, const Gain &gain);
  Gain pop();
  Gain top() const;
  void remove(const pair<int, int> &key);
  void remove(int part_id, int vertex_id);
  bool isActive() const;
  int total() const;
  void clear();

 private:
  MyPriorityQueueAllInOne bucket_;
};

#endif