TrickmanOff

Divide to k heaps in O(2^n * n)

Feb 13th, 2021
1,132
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
C++ 3.74 KB | None | 0 0
  1. //#pragma optimization_level 3
  2. //#pragma GCC target("sse,sse2,sse3,ssse3,sse4,popcnt,abm,mmx,avx,avx2,tune=native")
  3. //#pragma GCC optimize("Ofast,no-stack-protector,unroll-loops,fast-math")
  4.  
  5. #include <iostream>
  6. #include <algorithm>
  7. #include <vector>
  8. #include <queue>
  9. #include <functional>
  10. #include <set>
  11. #include <map>
  12. #include <math.h>
  13. #include <fstream>
  14. #include <cmath>
  15. #include <string>
  16. #include <random>
  17. #include <unordered_set>
  18. #include <unordered_map>
  19. #include <string.h>
  20. #include <stack>
  21. #include <assert.h>
  22. #include <list>
  23. #include <time.h>
  24. #include <memory>
  25. #include <chrono>
  26. using namespace std;
  27. //
  28. #define fast cin.tie(0);cout.tie(0);cin.sync_with_stdio(0);cout.sync_with_stdio(0);
  29. //#define cin in
  30. //#define cout out
  31. #define ll long long
  32. #define db double
  33. #define ld long double
  34. #define uset unordered_set
  35. #define umap unordered_map
  36. #define ms multiset
  37. #define pb push_back
  38. #define pq priority_queue
  39. #define umap unordered_map
  40. #define uset unordered_set
  41. #define ull unsigned long long
  42. #define pii pair<int, int>
  43. #define pll pair<ll, ll>
  44. #define pdd pair<ld, ld>
  45. #define pnn pair<Node*, Node*>
  46. #define uid uniform_int_distribution
  47. #define PI acos(-1.0)
  48. //#define sort(a, b) sort(a.begin(), a.end(), b())
  49. mt19937 rnd(chrono::steady_clock::now().time_since_epoch().count());
  50. ifstream in("input.txt");
  51. ofstream out("output.txt");
  52.  
  53. // divides balls to 'k' heaps in O(2^n * n)
  54.  
  55. vector<int> sum;
  56. vector<int> sm;
  57. vector<int> weight;
  58. int n, m_cnt, full_mask;
  59. int k;
  60.  
  61. bool is_bit(int a, int p) {
  62.     return (a >> p) % 2;
  63. }
  64.  
  65. void precalc_sums() {
  66.     for (int mask = 0; mask < m_cnt; ++mask) {
  67.         for (int bit = 0; bit < n; ++bit) {
  68.             if (is_bit(mask, bit))
  69.                 sum[mask] += weight[bit];
  70.         }
  71.     }
  72. }
  73.  
  74. void bfs() {
  75.     queue<int> q;
  76.     stack<int> nxt;
  77.    
  78.     int w = sum[full_mask] / k;
  79.    
  80.     for (int mask = 0; mask < m_cnt; ++mask) {
  81.         if (sum[mask] == w)
  82.             q.push(mask);
  83.     }
  84.    
  85.     for (int cur_w = w; cur_w <= sum[full_mask]; cur_w += w) {
  86.        
  87.         while (!q.empty() && sum[q.front()] < cur_w + w) {
  88.             int mask = q.front();
  89.             q.pop();
  90.            
  91.             for (int bit = 0; bit < n; ++bit) {
  92.                 if (is_bit(mask, bit)) continue;
  93.                 int n_mask = mask | (1 << bit);
  94.                
  95.                 if (sm[n_mask] || sum[n_mask] > cur_w + w) continue;
  96.                
  97.                 if (sum[mask] == cur_w)
  98.                     sm[n_mask] = mask;
  99.                 else if (sm[mask])
  100.                     sm[n_mask] = sm[mask];
  101.                
  102.                 if (sum[n_mask] == cur_w + w)
  103.                     nxt.push(n_mask);
  104.                 else
  105.                     q.push(n_mask);
  106.             }
  107.         }
  108.        
  109.         if (nxt.empty()) break;
  110.         while (!nxt.empty()) {
  111.             q.push(nxt.top());
  112.             nxt.pop();
  113.         }
  114.     }
  115. }
  116.  
  117. int main() {
  118.     cin >> n >> k;
  119.     weight.resize(n);
  120.     for (int& w : weight) cin >> w;
  121.     m_cnt = (1 << n);
  122.     sum.resize(m_cnt, 0);
  123.     sm.resize(m_cnt, 0);
  124.    
  125.     precalc_sums();
  126.     full_mask = (1 << n) - 1;
  127.    
  128.     if (sum[full_mask] % k) {
  129.         cout << "Impossible";
  130.     } else {
  131.         bfs();
  132.        
  133.         if (!sm[full_mask]) {
  134.             cout << "Impossible";
  135.             return 0;
  136.         }
  137.        
  138.         int w = sum[full_mask] / k;
  139.        
  140.         int mask = full_mask;
  141.         while (sum[mask] != 0) {
  142.             int cur = mask ^ sm[mask];
  143.             for (int bit = 0; bit < n; ++bit) {
  144.                 if (is_bit(cur, bit))
  145.                     cout << weight[bit] << ' ';
  146.             }
  147.             cout << '\n';
  148.            
  149.             mask = sm[mask];
  150.         }
  151.     }
  152. }
  153.  
Advertisement
Add Comment
Please, Sign In to add comment