library

This documentation is automatically generated by online-judge-tools/verification-helper

View the Project on GitHub maspypy/library

:warning: ds/offline_query/offline_set_intersection.hpp

Depends on

Code

#include "ds/csr.hpp"
#include "ds/to_small_key.hpp"

// given set S[0],...,S[N-1]
// Q query: calc |S[i] cap S[j]|
// M:=sum of size
// complexity: M sqrt{Q}/8
//
// https://codeforces.com/contest/2155/problem/F
// N,M,Q=300000, 300ms 程度
struct Offline_Set_Intersection {
  int N;
  To_Small_Key TSK;
  bool calculated;
  vc<pair<int, int>> dat;
  HashMap<int> query_id;
  vc<pair<int, int>> unique_query;
  vc<int> ids;

  // N: 集合の個数, K: 要素の種類数
  Offline_Set_Intersection(int N) : N(N), calculated(0) {}

  // x in S[i]
  // 同じ要素を 2 回登録すると壊れる(検査しない)
  void add(int i, int x) {
    assert(!calculated && 0 <= i && i < N);
    int k = TSK.query(x, true);
    dat.eb(i, k);
  }

  int get_qid(int i, int j) {
    if (i > j) swap(i, j);
    u64 k = u64(i) << 32 | u64(j);
    if (!query_id.count(k)) {
      query_id[k] = len(unique_query);
      unique_query.eb(i, j);
    }
    return query_id[k];
  }

  void query(int i, int j) {
    assert(!calculated && 0 <= i && i < N && 0 <= j && j < N);
    ids.eb(get_qid(i, j));
  }

  vc<int> calc() {
    assert(!calculated);
    calculated = true;
    int K = TSK.kind;
    int Q = len(unique_query);
    int B = sqrt(Q) / 8;
    vc<int> F(K);
    for (auto &[i, k] : dat) F[k]++;
    vc<int> heavy;
    FOR(k, K) if (F[k] >= B) heavy.eb(k);

    // StoX は light only
    CSR<int> StoX(N), XtoS(K), StoQ(N);
    for (auto &[i, k] : dat) {
      XtoS.add(k, i);
      if (F[k] < B) StoX.add(i, k);
    }
    FOR(q, Q) {
      auto [i, j] = unique_query[q];
      StoQ.add(i, q);
    }
    StoX.build(), XtoS.build(), StoQ.build();

    vc<int> ANS(Q);

    // heavy
    {
      vc<u64> A(N);
      vc<int> vis;
      for (int p = 0; p < len(heavy); p += 64) {
        vis.clear();
        // item [p,p+64)
        for (int idx = p; idx < p + 64; ++idx) {
          if (len(heavy) <= idx) break;
          for (auto &i : XtoS[heavy[idx]]) {
            A[i] ^= u64(1) << (idx - p);
            vis.eb(i);
          }
        }
        for (int q = 0; q < Q; ++q) {
          auto [i, j] = unique_query[q];
          ANS[q] += popcnt(A[i] & A[j]);
        }
        for (int i : vis) A[i] = 0;
      }
    }
    // light
    vc<int> A(N);
    FOR(i, N) {
      if (StoX[i].empty() || StoQ[i].empty()) continue;
      for (int x : StoX[i]) {
        for (int j : XtoS[x]) {
          A[j]++;
        }
      }
      for (int q : StoQ[i]) {
        int j = unique_query[q].se;
        ANS[q] += A[j];
      }
      for (int x : StoX[i]) {
        for (int j : XtoS[x]) {
          A[j]--;
        }
      }
    }
    ANS = rearrange(ANS, ids);
    return ANS;
  }
};
#line 1 "ds/csr.hpp"

template <typename T>
struct CSR {
  int n;
  bool prepared;
  vc<int> ptr;
  vc<int> I;
  vc<T> dat;

  CSR(int n = 0) : n(n), prepared(false) {}
  void add(int i, const T& x) {
    assert(0 <= i && i < n && !prepared);
    I.eb(i), dat.eb(x);
  }

  void build() {
    assert(!prepared);
    prepared = 1;
    ptr.assign(n + 1, 0);
    for (auto& i : I) ptr[1 + i]++;
    FOR(i, len(ptr) - 1) ptr[i + 1] += ptr[i];
    vc<T> tmp(len(dat));
    FOR(k, len(dat)) {
      int i = I[k];
      tmp[ptr[i]++] = dat[k];
    }
    swap(dat, tmp);
    ptr.pop_back();
    ptr.insert(ptr.begin(), 0);
    I.clear(), I.shrink_to_fit();
  }

  struct range {
    T *first, *last;
    T* begin() const { return first; }
    T* end() const { return last; }
    bool empty() const { return first == last; }
    int size() const { return last - first; }
  };

  range operator[](int i) {
    assert(prepared);
    return range{dat.data() + ptr[i], dat.data() + ptr[i + 1]};
  }
};
#line 2 "ds/hashmap.hpp"

// u64 -> Val

template <typename Val>
struct HashMap {
  // n は入れたいものの個数で ok

  HashMap(u32 n = 0) { build(n); }
  void build(u32 n) {
    u32 k = 8;
    while (k < n * 2) k *= 2;
    cap = k / 2, mask = k - 1;
    key.resize(k), val.resize(k), used.assign(k, 0);
  }

  // size を保ったまま. size=0 にするときは build すること.

  void clear() {
    used.assign(len(used), 0);
    cap = (mask + 1) / 2;
  }
  int size() { return len(used) / 2 - cap; }

  int index(const u64& k) {
    int i = 0;
    for (i = hash(k); used[i] && key[i] != k; i = (i + 1) & mask) {}
    return i;
  }

  Val& operator[](const u64& k) {
    if (cap == 0) extend();
    int i = index(k);
    if (!used[i]) { used[i] = 1, key[i] = k, val[i] = Val{}, --cap; }
    return val[i];
  }

  Val get(const u64& k, Val default_value) {
    int i = index(k);
    return (used[i] ? val[i] : default_value);
  }

  bool count(const u64& k) {
    int i = index(k);
    return used[i] && key[i] == k;
  }

  // f(key, val)

  template <typename F>
  void enumerate_all(F f) {
    FOR(i, len(used)) if (used[i]) f(key[i], val[i]);
  }

private:
  u32 cap, mask;
  vc<u64> key;
  vc<Val> val;
  vc<bool> used;

  u64 hash(u64 x) {
    static const u64 FIXED_RANDOM = std::chrono::steady_clock::now().time_since_epoch().count();
    x += FIXED_RANDOM;
    x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9;
    x = (x ^ (x >> 27)) * 0x94d049bb133111eb;
    return (x ^ (x >> 31)) & mask;
  }

  void extend() {
    vc<pair<u64, Val>> dat;
    dat.reserve(len(used) / 2 - cap);
    FOR(i, len(used)) {
      if (used[i]) dat.eb(key[i], val[i]);
    }
    build(2 * len(dat));
    for (auto& [a, b]: dat) (*this)[a] = b;
  }
};
#line 2 "ds/to_small_key.hpp"

// [30,10,20,30] -> [0,1,2,0] etc.
struct To_Small_Key {
  int kind = 0;
  HashMap<int> MP;
  vc<u64> raw;
  To_Small_Key(u32 n = 0) : MP(n) {}
  void reserve(u32 n) { MP.build(n); }
  int size() { return MP.size(); }
  u64 restore(int i) { return raw[i]; }
  int query(u64 x, bool set_if_not_exist) {
    int ans = MP.get(x, -1);
    if (ans == -1 && set_if_not_exist) {
      raw.eb(x);
      MP[x] = ans = kind++;
    }
    return ans;
  }
};
#line 3 "ds/offline_query/offline_set_intersection.hpp"

// given set S[0],...,S[N-1]
// Q query: calc |S[i] cap S[j]|
// M:=sum of size
// complexity: M sqrt{Q}/8
//
// https://codeforces.com/contest/2155/problem/F
// N,M,Q=300000, 300ms 程度
struct Offline_Set_Intersection {
  int N;
  To_Small_Key TSK;
  bool calculated;
  vc<pair<int, int>> dat;
  HashMap<int> query_id;
  vc<pair<int, int>> unique_query;
  vc<int> ids;

  // N: 集合の個数, K: 要素の種類数
  Offline_Set_Intersection(int N) : N(N), calculated(0) {}

  // x in S[i]
  // 同じ要素を 2 回登録すると壊れる(検査しない)
  void add(int i, int x) {
    assert(!calculated && 0 <= i && i < N);
    int k = TSK.query(x, true);
    dat.eb(i, k);
  }

  int get_qid(int i, int j) {
    if (i > j) swap(i, j);
    u64 k = u64(i) << 32 | u64(j);
    if (!query_id.count(k)) {
      query_id[k] = len(unique_query);
      unique_query.eb(i, j);
    }
    return query_id[k];
  }

  void query(int i, int j) {
    assert(!calculated && 0 <= i && i < N && 0 <= j && j < N);
    ids.eb(get_qid(i, j));
  }

  vc<int> calc() {
    assert(!calculated);
    calculated = true;
    int K = TSK.kind;
    int Q = len(unique_query);
    int B = sqrt(Q) / 8;
    vc<int> F(K);
    for (auto &[i, k] : dat) F[k]++;
    vc<int> heavy;
    FOR(k, K) if (F[k] >= B) heavy.eb(k);

    // StoX は light only
    CSR<int> StoX(N), XtoS(K), StoQ(N);
    for (auto &[i, k] : dat) {
      XtoS.add(k, i);
      if (F[k] < B) StoX.add(i, k);
    }
    FOR(q, Q) {
      auto [i, j] = unique_query[q];
      StoQ.add(i, q);
    }
    StoX.build(), XtoS.build(), StoQ.build();

    vc<int> ANS(Q);

    // heavy
    {
      vc<u64> A(N);
      vc<int> vis;
      for (int p = 0; p < len(heavy); p += 64) {
        vis.clear();
        // item [p,p+64)
        for (int idx = p; idx < p + 64; ++idx) {
          if (len(heavy) <= idx) break;
          for (auto &i : XtoS[heavy[idx]]) {
            A[i] ^= u64(1) << (idx - p);
            vis.eb(i);
          }
        }
        for (int q = 0; q < Q; ++q) {
          auto [i, j] = unique_query[q];
          ANS[q] += popcnt(A[i] & A[j]);
        }
        for (int i : vis) A[i] = 0;
      }
    }
    // light
    vc<int> A(N);
    FOR(i, N) {
      if (StoX[i].empty() || StoQ[i].empty()) continue;
      for (int x : StoX[i]) {
        for (int j : XtoS[x]) {
          A[j]++;
        }
      }
      for (int q : StoQ[i]) {
        int j = unique_query[q].se;
        ANS[q] += A[j];
      }
      for (int x : StoX[i]) {
        for (int j : XtoS[x]) {
          A[j]--;
        }
      }
    }
    ANS = rearrange(ANS, ids);
    return ANS;
  }
};
Back to top page