Program Listing for File defs.h

Return to documentation for file (src/partition/defs.h)

// 超图分割专用数据结构定义
#ifndef DEFS_H
#define DEFS_H

#include <tbb/enumerable_thread_specific.h>
#include <tbb/global_control.h>
#include <tbb/parallel_for.h>
#include <tbb/parallel_sort.h>

#include <algorithm>
#include <atomic>
#include <cassert>
#include <cfloat>
#include <eigen3/Eigen/Dense>
#include <eigen3/Eigen/Sparse>
#include <fstream>
#include <functional>
#include <iomanip>
#include <iostream>
#include <limits>
#include <memory>
#include <mutex>
#include <numeric>
#include <queue>
#include <random>
#include <set>
#include <shared_mutex>
#include <sstream>
#include <stdexcept>
#include <string>
#include <type_traits>
#include <utility>
#include <variant>
#include <vector>

#include "../sta/HSIOTdmDelay.h"
#include "../sta/HSTimingSlackRange.h"
#include "json.hpp"
#include "parallel_hashmap/phmap.h"
#include "spdlog/spdlog.h"


using namespace std;
using namespace Eigen;
using phmap::flat_hash_map;
using json = nlohmann::json;

struct curPartStatus {
  // 不变值:FPGA级和Die级的数据相同
  float clock_period_max = -numeric_limits<float>::max();
  float original_slack_max = numeric_limits<float>::lowest();
  float original_slack_min = numeric_limits<float>::max();

  // 变化值:FPGA级和Die级的数据不同
  flat_hash_map<float, float> clock_period_delays;
};

struct cpAttr {
  int num;
  float normalize_delay_cut;
  float normalize_delay_cut_die;
};

struct mulClockAttr {
  // 共性值
  float original_delay_cut_factor;
  float original_delay_cut;
  float original_delay_cut_die;

  flat_hash_map<float, shared_ptr<cpAttr>> clock_periods;

  // 状态值(差异值)
  shared_ptr<curPartStatus> status;
};

struct PartitionConstraintConfig {
  bool has_fix = false;
  bool has_region =
      false;
  bool has_io =
      false;
  bool has_timing =
      false;
  bool force_topo =
      false;
  bool has_timing_base =
      false;
  bool has_io_base =
      false;
  int max_hop = -1;
};
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE_WITH_DEFAULT(PartitionConstraintConfig,
                                                has_fix, has_region, has_io,
                                                has_timing, force_topo,
                                                has_timing_base, has_io_base,
                                                max_hop)


struct CoarsenConfig {
  int level = 30;
  int large_net_threshold =
      200;
  int thr_coarsen_vertice =
      30;
  double coarsening_ratio =
      1.6;
  int coarsen_method =
      1;
  int enable_weight_penalty =
      0;
};
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE_WITH_DEFAULT(
    CoarsenConfig, level, large_net_threshold, thr_coarsen_vertice,
    coarsening_ratio, coarsen_method, enable_weight_penalty)


struct InitialPartitionConfig {
  int num_initial_solutions = 64;
};
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE_WITH_DEFAULT(InitialPartitionConfig,
                                                num_initial_solutions)


struct RefineStageConfig {
  int max_move = 1000;
  int max_neg_move = 10;
  int refine_iters = 1;
  int num_best_initial_solutions = 8;
  int large_net_threshold = 200;

  bool enable_Greedy = true;
  bool enable_HER = false;
  bool enable_PM = true;
  bool enable_FM = false;
  bool enable_DSFM_S = false;
  bool enable_DSFM_M = false;
  int enable_mixed_refinement =
      0;
  bool enable_parallel_refine = false;
};
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE_WITH_DEFAULT(
    RefineStageConfig, max_move, max_neg_move, refine_iters,
    num_best_initial_solutions, large_net_threshold, enable_Greedy, enable_HER,
    enable_PM, enable_FM, enable_DSFM_S, enable_DSFM_M, enable_mixed_refinement,
    enable_parallel_refine)


struct PartitionFlowConfig {
  int mode = 0;
  int COCP_on = 1;
  int V_Cycle_on        = 1;
  int V_Cycle_run       = 2;
  int hold_solution_num = 4;
};
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE_WITH_DEFAULT(PartitionFlowConfig, mode,
                                                COCP_on, V_Cycle_on,
                                                V_Cycle_run, hold_solution_num)


struct TimingPartitionConfig {
  bool skip_path_timing =
      false;
  bool has_mul_clock_domains =
      false;

  int timing_hold_solution_num = 1;
  int guardband_flag = 1;

  // 延迟模型参数
  float extra_delay_cut_factor = 1.0f;
  float extra_delay_cut = 100.0f;
  float extra_delay_cut_die =
      2.0f;
  float extra_delay_cut_tdm = 52.657f;
  float extra_delay_tdm = 0.5993f;

  // 时序代价权重因子
  float net_timing_factor_coarse = 0.0f;
  float net_timing_factor_refine = 1.0f;
  float path_timing_factor = 1.0f;
  float path_snaking_factor = 1.0f;
  float timing_exp_factor =
      2.0f;
  float clock_period = 0.0f;

  // slack 传播与更新策略
  int has_timing_propogate =
      1;
  int propogate_effort_value = 5;
  int path_slack_update_mode =
      1;
  int enable_partition_HSFulltiming =
      1;
  int enablePhase2Verification =
      0;
  int only_compare_timing =
      0;

  shared_ptr<mulClockAttr> mul_clock_attr;
};
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE_WITH_DEFAULT(
    TimingPartitionConfig, skip_path_timing, has_mul_clock_domains,
    timing_hold_solution_num, guardband_flag, extra_delay_cut_factor,
    extra_delay_cut, extra_delay_cut_die, extra_delay_cut_tdm, extra_delay_tdm,
    net_timing_factor_coarse, net_timing_factor_refine, path_timing_factor,
    path_snaking_factor, timing_exp_factor, clock_period, has_timing_propogate,
    propogate_effort_value, path_slack_update_mode,
    enable_partition_HSFulltiming, enablePhase2Verification,
    only_compare_timing)


struct PartitionParams {
  PartitionFlowConfig flow;
  PartitionConstraintConfig constraints;
  CoarsenConfig coarsen;
  InitialPartitionConfig initial;
  RefineStageConfig refine;
  TimingPartitionConfig timing;
};
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE_WITH_DEFAULT(PartitionParams, flow,
                                                constraints, coarsen, initial,
                                                refine, timing)


struct TDMParams {
  int max_iters                = 100;   // 最大迭代次数
  double convergence_threshold = 1e-3;  // 提前终止阈值
  double convergence_num       = 10;    // 提前终止次数
  double stable_num            = 3;     // 稳定次数
  double decay_factor          = 0.7;   // 控制 Lagrange下降速率
  double initial_learning_rate = 0.2;  // 初始步长系数,控制初始更新速度
  double decay_base            = 0.5;    // 指数衰减基数
  double decay_rate            = 0.02;   // 衰减率
  bool enable_net              = false;  // 是否开启分析net的时序信息
  bool avg_only = false;  // 是否开启“仅按割边数/容量均分”的非时序TDM
  bool tdm_fast_mode = false;  // 快速模式: 仅评估精简策略集合
  bool use_worst_slack_objective =
      false;  // 连续阶段目标: false=arrival_time, true=worst_slack
  bool allow_mgt_on_non_timing_edges =
      false;  // 允许 MGT 用于不在 cut_timing_path 上的 is_tdm 边
  string tdm_opt_mode =
      "legacy";  // legacy/root_split/final_polish/v2/no_postprocess
  string delay_lib_path;                  // 延迟库文件路径
  int tdm_topk                      = 1;  // 输出Top K条时序路径
  int gio_channel_grouping_capacity = 23;  // 单根GIO可承载的逻辑通道数
};

struct RoutingParams {
  double routing_cost_factor = 0.5;
  string timing_routing_cost_mode      = "critical_minmax_hop";
  double timing_hop_penalty_multiplier = 0.5;
};

struct Metrics {
  int cut = 0;
  int tdm = 0;
  int topo = 0;
  int violation = 0;
  double cost = 0.0;
  double worstSlack = 0.0;
  int maxHop = 0;
  float tdmDelay = 0.0;
  float topkratio = 0.0;
};

struct HierarchyPartitionCutRecord {
  vector<int> path_prefix;
  int depth = 0;
  int part_count = 0;
  int node_count = 0;
  int net_count = 0;
  vector<int> part_node_counts;
  Metrics metric;
  vector<vector<int>> cut_weights;
};

struct SimpleTiming {
  vector<int> path_node;
  vector<int> path_net;

  string slack;
  string period;
  string clock_domain;
  int pathId;

  SimpleTiming() = default;

  SimpleTiming(const vector<int> &path_node_arg,
               const vector<int> &path_net_arg, string slack_arg,
               string period_arg, string clock_domain_arg, int pathId_arg)
  {
    path_node    = path_node_arg;
    path_net     = path_net_arg;
    slack        = slack_arg;
    period       = period_arg;
    clock_domain = clock_domain_arg;
    pathId       = pathId_arg;
  }
};

struct TimingPath {
  vector<int> path;
  vector<int> arcs;

  // 相对值: 设置guardband后的&归一化后的值
  float slack = 0.0;
  float slack_new = 0.0;

  // 绝对值
  float cp = 0.0;
  float slack_old = 0.0;
  float slack_new_abolute = 0.0;

  int cut = 0;
  int dHop = 0;
  int index = 0;
  bool is_path_in_cell = false;

  TimingPath() = default;

  TimingPath(const vector<int> &path_arg, const vector<int> &arcs_arg,
             float slack_arg, float cp_arg)
  {
    path  = path_arg;
    arcs  = arcs_arg;
    slack = slack_arg;
    cp    = cp_arg;
    // 初始状态认为还未产生新的最优化结果,设定为上限不可及的极大值
    slack_new = numeric_limits<float>::max();
  }
};

class node {
 public:
  int weight = 1;
  int begin = 0;
  int size = 0;
  VectorXi resources;

  bool operator>(const node &n) const
  {
    int gap = (resources - n.resources).minCoeff();
    return gap >= 0;
  }

  void add(const node &n)
  {
    resources += n.resources;
  }

  void substract(const node &n)
  {
    resources -= n.resources;
  }
};

class net {
 public:
  double weight;

  int begin = 0;
  int size = 0;

  bool operator<(const net &p) const
  {
    return weight > p.weight || (weight == p.weight && size < p.size);
  }

  static bool compareByNodeSize(const net &a, const net &b)
  {
    return a.size < b.size;
  }

  static bool comparePairs(const pair<int, int> &a, const pair<int, int> &b)
  {
    if (a.first != b.first) {
      // 按照第一个元素从小到大排序
      return a.first < b.first;
    } else {
      // 对于第一个元素相等的,按照第二个元素从大到小排序
      return a.second > b.second;
    }
  }
};

struct pinSlack {
  float slack;
  float cp;
  pair<int, int> pinPreId;
  pair<int, int> pinPostId;

  friend void to_json(json &j, const pinSlack &ps);
  friend void from_json(const json &j, pinSlack &ps);
};

class timing {
 public:
  // slack information
  vector<TimingPath> timing_paths;
  flat_hash_map<pair<int, int>, pinSlack> pin_minslack;
  vector<float> hyperedge_slacks;

  // timing cost information
  vector<float> path_timing_cost;
  vector<float> hyperedge_timing_cost;

  // timing grapth related information
  vector<set<int>> ins_paths;
  flat_hash_map<int, set<pair<int, int>>> net_insPin_set;
};

class range {
 public:
  using iterator = std::vector<int>::const_iterator;

  range(iterator begin, iterator end) : begin_(begin), end_(end)
  {
  }

  iterator begin() const
  {
    return begin_;
  }

  iterator end() const
  {
    return end_;
  }

 private:
  iterator begin_;
  iterator end_;
};

class graph {
 public:
  vector<node> nodes;
  vector<net> nets;
  vector<bool> net_bypass;
  vector<int> fixed_assign;
  vector<bool> region_fixed;

  vector<int> incident_nodes;
  vector<int> incident_nets;
  vector<set<int>> candidates;
  vector<int> community;
  vector<VectorXd> embeddings;
  timing timingInfo;

  range get_incident_nodes(int net_id) const
  {
    return range(
        incident_nodes.begin() + nets[net_id].begin,
        incident_nodes.begin() + nets[net_id].begin + nets[net_id].size);
  }

  range get_incident_nets(int node_id) const
  {
    return range(
        incident_nets.begin() + nodes[node_id].begin,
        incident_nets.begin() + nodes[node_id].begin + nodes[node_id].size);
  }
};

struct ChannelStatus {
  int total_channels;
  vector<int> failed_channels;
  string errorno;
};

class DieConnection {
 public:
  int left_fpga;
  int right_fpga;
  int left_die;
  int right_die;
  string left_socket;
  string right_socket;
  bool isMGT;
  ChannelStatus left2right;
  ChannelStatus right2left;
};

class fpga {
 public:
  vector<VectorXi> resources;
  vector<vector<int>> topology;
  vector<int> maxDist;
  vector<vector<int>> dist;
  vector<vector<set<int>>> S_hat;
  vector<vector<int>> cutweights_assignment;
  vector<vector<int>> hio_channel_assignment;
  vector<vector<int>> mgt_channel_assignment;
  vector<vector<int>> fpga_hierarchy_paths;
  bool bound_constraint = false;
  bool need_routing = false;
  bool is_hierarchical = false;
  VectorXi upper_resources;
  VectorXi lower_resources;
  vector<DieConnection> die_connections;
  map<int, int> die_to_fpga_map;  // 存储 die 到 FPGA 的映射
  // flat_hash_map<int, int> die_to_fpga_map;  // 存储 die 到 FPGA 的映射
  int dies_per_fpga  = 1;
  int board_count    = -1;
  int fpga_per_board = -1;
  HSFullTiming::HSIOTdmDelay tdmEstimate;

  int computeCutWeight(const graph &g, const vector<int> &parts) const
  {
    int cutweight = 0, fpga_num = resources.size();
    for (int i = 0; i < g.nets.size(); i++) {
      vector<int> indicators(fpga_num, 0);
      auto iter = g.incident_nodes.begin() + g.nets[i].begin;
      for (auto j = g.incident_nodes.begin() + g.nets[i].begin + 1;
           j < g.incident_nodes.begin() + g.nets[i].begin + g.nets[i].size;
           j++) {
        indicators[parts[*j]] = 1;
      }
      for (int k = 0; k < fpga_num; k++) {
        if (indicators[k] == 1 && k != parts[*iter]) {
          cutweight += g.nets[i].weight;
        }
      }
    }
    return cutweight;
  }

  float computeMaxTdmDelayFromCutWeights(
      const vector<vector<int>> &cut_weights) const
  {
    if (tdmEstimate.isEmptyIO()) {
      return 0.0f;
    }

    auto tdm_estimate = tdmEstimate;
    tdm_estimate.setCutSizeArray(cut_weights);
    tdm_estimate.setupTdmCutDelay();

    if (tdm_estimate.isIllegal()) {
      return numeric_limits<float>::max();
    }

    float max_tdm_delay = 0.0f;
    const int fpga_num  = cut_weights.size();
    for (int i = 0; i < fpga_num; i++) {
      for (int j = i + 1; j < fpga_num; j++) {
        if (cut_weights[i][j] > 0) {
          max_tdm_delay = max(max_tdm_delay, tdm_estimate.getTdmCutDelay(i, j));
        }
      }
    }
    return max_tdm_delay;
  }

  Metrics computeCutWeightsOld(const graph &g, const vector<int> &parts,
                               vector<vector<int>> &cut_weights) const
  {
    int net_num = g.nets.size(), fpga_num = resources.size();
    fill(cut_weights.begin(), cut_weights.end(), vector<int>(fpga_num, 0));
    Metrics metric;
    metric.cut = metric.tdm = metric.topo = metric.violation = metric.tdmDelay =
        0;
    cut_weights.clear();
    cut_weights.resize(fpga_num);
    for (int i = 0; i < fpga_num; i++) {
      cut_weights[i].resize(fpga_num, 0);
    }
    vector<int> indicators(fpga_num, 0);
    for (int i = 0; i < net_num; i++) {
      fill(indicators.begin(), indicators.end(), 0);
      auto iter = g.incident_nodes.begin() + g.nets[i].begin;
      for (auto j = g.incident_nodes.begin() + g.nets[i].begin + 1;
           j < g.incident_nodes.begin() + g.nets[i].begin + g.nets[i].size;
           j++) {
        indicators[parts[*j]] = 1;
      }
      for (int k = 0; k < fpga_num; k++) {
        if (indicators[k] == 1 && k != parts[*iter]) {
          metric.cut += g.nets[i].weight;
          cut_weights[k][parts[*iter]] += g.nets[i].weight;
          cut_weights[parts[*iter]][k] += g.nets[i].weight;
        }
      }
    }
    for (size_t i = 0; i < cutweights_assignment.size(); ++i) {
      for (size_t j = i + 1; j < cutweights_assignment.size(); ++j) {
        if (cutweights_assignment[i][j] == 0) {
          metric.topo += cut_weights[i][j] * (dist[i][j] - 1);
          metric.violation += cut_weights[i][j];
        } else {
          metric.tdm =
              max(metric.tdm,
                  static_cast<int>(
                      ceil(static_cast<double>(cut_weights[i][j]) /
                           static_cast<double>(cutweights_assignment[i][j]))));
        }
      }
    }
    return metric;
  }

  Metrics computeCutWeights(const graph &g, const vector<int> &parts,
                            vector<vector<int>> &cut_weights,
                            bool withNet = false)
  {
    int fpga_num = cut_weights.size();
    int net_num  = g.nets.size();
    fill(cut_weights.begin(), cut_weights.end(), vector<int>(fpga_num, 0));
    Metrics metric;
    metric.cut = metric.tdm = metric.topo = metric.violation = metric.tdmDelay =
        0;
    vector<int> indicators(fpga_num, 0);
    for (int i = 0; i < net_num; i++) {
      fill(indicators.begin(), indicators.end(), 0);
      auto iter = g.incident_nodes.begin() + g.nets[i].begin;
      for (auto j = g.incident_nodes.begin() + g.nets[i].begin + 1;
           j < g.incident_nodes.begin() + g.nets[i].begin + g.nets[i].size;
           j++) {
        indicators[parts[*j]] = 1;
      }
      for (int k = 0; k < fpga_num; k++) {
        if (indicators[k] == 1 && k != parts[*iter]) {
          metric.cut += g.nets[i].weight;
          cut_weights[k][parts[*iter]] += g.nets[i].weight;
          cut_weights[parts[*iter]][k] += g.nets[i].weight;
        }
      }
    }
    const vector<vector<int>> logical_cut_weights = cut_weights;
    if (!tdmEstimate.isEmptyIO()) {
      if (withNet) {
        vector<HSFullTiming::PTNet> &netAll = tdmEstimate.getNetAll();
        netAll.clear();
        for (int i = 0; i < g.nets.size(); i++) {
          HSFullTiming::PTNet net;
          net.m_netNum = static_cast<unsigned short>(g.nets[i].weight);
          fill(indicators.begin(), indicators.end(), 0);
          auto iter = g.incident_nodes.begin() + g.nets[i].begin;
          net.m_nodes.push_back(parts[*iter]);
          for (auto j = g.incident_nodes.begin() + g.nets[i].begin + 1;
               j < g.incident_nodes.begin() + g.nets[i].begin + g.nets[i].size;
               j++) {
            indicators[parts[*j]] = 1;
          }
          for (int k = 0; k < fpga_num; k++) {
            if (indicators[k] == 1 && k != parts[*iter]) {
              net.m_nodes.push_back(k);
            }
          }
          if (net.m_nodes.size() > 1) {
            netAll.push_back(net);
          }
        }
      }
      tdmEstimate.setCutSizeArray(cut_weights);
      tdmEstimate.setupTdmCutDelay();
      cut_weights = tdmEstimate.getCutSizeArrayRouting();
      // metric.cut  = 0;
      for (size_t i = 0; i < fpga_num; ++i) {
        for (size_t j = i + 1; j < fpga_num; ++j) {
          // metric.cut += cut_weights[i][j];
          if (logical_cut_weights[i][j] > 0) {
            metric.tdmDelay =
                max(metric.tdmDelay, tdmEstimate.getTdmCutDelay(i, j));
          }
        }
      }
    }
    return metric;
  }
};

enum class HierarchyType { RACK, CLUSTER, BOARD, FPGA, DIE };

struct HierarchyLinkEndpoint {
  int node_id = -1;
  int die_id  = -1;
  string socket;
  string fpga_alias;
};

struct HierarchyLink {
  int link_id = -1;
  HierarchyLinkEndpoint left;
  HierarchyLinkEndpoint right;
  bool is_mgt = false;
  int channel_capacity = 0;
  int hio_channels = 0;
  int mgt_channels = 0;
  string cable;
};

class hierarchy {
 public:
  HierarchyType type;
  fpga fpgas0;
  fpga fpgas;
  vector<hierarchy> contents;
  int node_id        = -1;
  int parent_node_id = -1;
  string alias;
  vector<int> children_node_ids;
  flat_hash_map<string, int> alias_to_node_id;
  vector<HierarchyLink> topology_links;
};

class fixInfo {
 public:
  int fpgaNo;
  vector<vector<int>> id;
};

struct pq {
  int x, y;
  int z;
  bool operator<(const pq &rhs) const
  {
    if (x < rhs.x)
      return true;
    else if (x == rhs.x)
      if (z > rhs.z)
        return true;
      else if (z == rhs.z)
        return y < rhs.y;
      else
        return false;
    else
      return false;
  }
};

class DisjointSet {
  vector<int> parent;
  int size;

 public:
  DisjointSet(int size) : size(size)
  {
    parent.resize(size);
    for (int i = 0; i < size; i++) {
      parent[i] = i;
    }
  }
  int size_()
  {
    return size;
  }
  int find(int x)
  {
    while (x != parent[x]) {
      parent[x] = parent[parent[x]];
      x         = parent[x];
    }
    return x;
  }
  void union_(int x, int y)
  {
    parent[find(y)] = find(x);
  }
  int list(vector<int> &result)
  {
    flat_hash_map<int, int> visited;
    result.resize(size);
    int index = 0;
    for (int i = 0; i < size; i++) {
      int ancestor = find(i);
      if (visited.find(ancestor) != visited.end()) {
        result[i] = visited[ancestor];
      } else {
        visited[ancestor] = index;
        result[i]         = index;
        index++;
      }
    }
    return index;
  }
};

class comb_set  // 并查集
{
 public:
  comb_set(int n);
  void merge(int x, int y);
  int getfa(int x);

 private:
  vector<int> fa;
};

namespace kahypar {
template <typename Key, typename Value, typename Derived>
class SparseMapBase {
 protected:
  struct MapElement {
    Key key;
    Value value;
  };

 public:
  SparseMapBase(const SparseMapBase &)            = delete;
  SparseMapBase &operator=(const SparseMapBase &) = delete;

  SparseMapBase &operator=(SparseMapBase &&) = delete;

  size_t size() const
  {
    return _size;
  }

  bool contains(const Key key) const
  {
    return static_cast<const Derived *>(this)->containsImpl(key);
  }

  void add(const Key key, const Value value)
  {
    static_cast<Derived *>(this)->addImpl(key, value);
  }

  const MapElement *begin() const
  {
    return _dense;
  }

  const MapElement *end() const
  {
    return _dense + _size;
  }

  MapElement *begin()
  {
    return _dense;
  }

  MapElement *end()
  {
    return _dense + _size;
  }

  void clear()
  {
    static_cast<Derived *>(this)->clearImpl();
  }

  Value &operator[](const Key key)
  {
    const size_t index = _sparse[key];
    if (!contains(key)) {
      _dense[_size] = MapElement{key, Value()};
      _sparse[key]  = _size++;
      return _dense[_size - 1].value;
    }
    return _dense[index].value;
  }

  const Value &get(const Key key) const
  {
    return _dense[_sparse[key]].value;
  }

 protected:
  explicit SparseMapBase(const size_t max_size, const Value initial_value = 0)
      : _size(0),
        _sparse(std::make_unique<size_t[]>(
            (max_size * sizeof(MapElement) + max_size * sizeof(size_t)) /
            sizeof(size_t))),
        _dense(nullptr)
  {
    _dense = reinterpret_cast<MapElement *>(_sparse.get() + max_size);
    for (size_t i = 0; i < max_size; ++i) {
      _sparse[i] = std::numeric_limits<size_t>::max();
      _dense[i]  = MapElement{std::numeric_limits<Key>::max(), initial_value};
    }
  }

  ~SparseMapBase() = default;

  SparseMapBase(SparseMapBase &&other)
      : _size(other._size),
        _sparse(std::move(other._sparse)),
        _dense(std::move(other._dense))
  {
    other._size   = 0;
    other._sparse = nullptr;
    other._dense  = nullptr;
  }

  size_t _size;
  std::unique_ptr<size_t[]> _sparse;
  MapElement *_dense;
};

template <typename Key, typename Value>
class SparseMap final
    : public SparseMapBase<Key, Value, SparseMap<Key, Value>> {
  using Base = SparseMapBase<Key, Value, SparseMap<Key, Value>>;
  friend Base;

 public:
  explicit SparseMap(const Key max_size, const Value initial_value = 0)
      : Base(max_size, initial_value)
  {
  }

  SparseMap(const SparseMap &)                 = delete;
  SparseMap &operator=(const SparseMap &other) = delete;

  SparseMap(SparseMap &&other) : Base(std::move(other))
  {
  }

  SparseMap &operator=(SparseMap &&other)
  {
    _sparse       = std::move(other._sparse);
    _size         = 0;
    _dense        = std::move(other._dense);
    other._size   = 0;
    other._sparse = nullptr;
    other._dense  = nullptr;
    return *this;
  }

  ~SparseMap() = default;

  void remove(const Key key)
  {
    const size_t index = _sparse[key];
    if (index < _size && _dense[index].key == key) {
      std::swap(_dense[index], _dense[_size - 1]);
      _sparse[_dense[index].key] = index;
      --_size;
    }
  }

 private:
  bool containsImpl(const Key key) const
  {
    const size_t index = _sparse[key];
    return index < _size && _dense[index].key == key;
  }

  void addImpl(const Key key, const Value value)
  {
    const size_t index = _sparse[key];
    if (index >= _size || _dense[index].key != key) {
      _dense[_size] = {key, value};
      _sparse[key]  = _size++;
    }
  }

  void clearImpl()
  {
    _size = 0;
  }

  using Base::_dense;
  using Base::_size;
  using Base::_sparse;
};
};  // namespace kahypar

class SpinLock {
 public:
  // boilerplate to make it 'copyable'. but we just clear the spinlock. there is
  // never a use case to copy a locked spinlock
  SpinLock()
  {
  }
  SpinLock(const SpinLock &)
  {
  }
  SpinLock &operator=(const SpinLock &)
  {
    spinner.clear(std::memory_order_relaxed);
    return *this;
  }

  bool tryLock()
  {
    return !spinner.test_and_set(std::memory_order_acquire);
  }

  void lock()
  {
    while (spinner.test_and_set(std::memory_order_acquire)) {
      // spin
      // stack overflow says adding 'cpu_relax' instruction may improve
      // performance
    }
  }

  void unlock()
  {
    spinner.clear(std::memory_order_release);
  }

 private:
  std::atomic_flag spinner = ATOMIC_FLAG_INIT;
};

template <typename Value>
class ConcurrentBucketMap {
  static constexpr size_t BUCKET_FACTOR = 128;

  using Bucket = std::vector<Value>;

 public:
  ConcurrentBucketMap()
      : _num_buckets(align_to_next_power_of_two(
            BUCKET_FACTOR * std::thread::hardware_concurrency())),
        _mod_mask(_num_buckets - 1),
        _spin_locks(_num_buckets),
        _buckets(_num_buckets)
  {
  }

  ConcurrentBucketMap(const ConcurrentBucketMap &)            = delete;
  ConcurrentBucketMap &operator=(const ConcurrentBucketMap &) = delete;

  ConcurrentBucketMap(ConcurrentBucketMap &&other)
      : _num_buckets(other._num_buckets),
        _mod_mask(_num_buckets - 1),
        _spin_locks(_num_buckets),
        _buckets(std::move(other._buffer))
  {
  }

  template <typename F>
  void doParallelForAllBuckets(const F &f)
  {
    tbb::parallel_for(static_cast<size_t>(0), _num_buckets,
                      [&](const size_t i) { f(i); });
  }

  // ! Returns the number of buckets
  size_t numBuckets() const
  {
    return _num_buckets;
  }

  // ! Returns the corresponding bucket
  Bucket &getBucket(const size_t bucket)
  {
    return _buckets[bucket];
  }

  // ! Reserves memory in each bucket such that the estimated number of
  // insertions ! can be handled without the need (with high probability) of
  // expensive bucket resizing.
  void reserve_for_estimated_number_of_insertions(
      const size_t estimated_num_insertions)
  {
    // ! Assumption is that keys are evenly distributed among buckets (with a
    // small buffer)
    const size_t estimated_bucket_size = std::max(
        static_cast<size_t>(1.5 * estimated_num_insertions) / _num_buckets,
        static_cast<size_t>(1));
    tbb::parallel_for(
        static_cast<size_t>(0), _num_buckets,
        [&](const size_t i) { _buckets[i].reserve(estimated_bucket_size); });
  }

  // ! Inserts a key-value pair
  void insert(const size_t &key, Value &&value)
  {
    size_t bucket = key & _mod_mask;
    _spin_locks[bucket].lock();
    _buckets[bucket].emplace_back(std::move(value));
    _spin_locks[bucket].unlock();
  }

  // ! Frees the memory of all buckets
  void free()
  {
    tbb::parallel_for(static_cast<size_t>(0), _buckets.size(),
                      [&](const size_t i) {
                        std::vector<Value> tmp_vec;
                        _buckets[i] = std::move(tmp_vec);
                      });
  }

  // ! Frees the memory of the corresponding bucket
  void free(const size_t bucket)
  {
    std::vector<Value> tmp_vec;
    _buckets[bucket] = std::move(tmp_vec);
  }

  // ! Clears the corresponding bucket
  void clear(const size_t bucket)
  {
    _buckets[bucket].clear();
  }

  void clearParallel()
  {
    doParallelForAllBuckets([&](const size_t i) { clear(i); });
  }

 private:
  size_t align_to_next_power_of_two(const size_t size) const
  {
    return std::pow(2.0, std::ceil(std::log2(static_cast<double>(size))));
  }

  const size_t _num_buckets;
  const size_t _mod_mask;
  std::vector<SpinLock> _spin_locks;
  std::vector<Bucket> _buckets;
};

#endif