CSU-ACM2025 暑期训练赛-第三场 题解
Problem A - Fast XORting
2023 ICPC Southeastern Europe Regional Contest F
\(n\) 为 \(2\) 的幂,要将 \(0\) 到 \(n-1\) 这 \(n\) 个数升序排列,有两种操作:
- 操作一:将相邻两个数交换
- 操作二:将整个数组对x取异或
求最小操作次数。
首先,如果不进行操作二,那么操作一的次数显然等于逆序对的个数,逆序对个数可以用树状数组\(O(n\log n)\)实现。
接下来考虑进行操作二的情况。可以发现操作二与操作一的顺序并不影响最终结果,所以我们可以认为先进行操作二再进行操作一。且两次操作二总能转化为一次操作二,也就是说操作二最多进行一次。
现在问题转化为,进行一次操作二,使逆序对个数最小化。
我们可以从小到大枚举 \(x\) 的二进制下每一位,判断使逆序对最小化需要该位为 \(0\) 还是为 \(1\),位之间是互不影响的,最终的时间复杂度为 \(O(n\log^2 n)\)。
接下来证明位之间互不影响:
假设 \(a=2^x,b=2^y,x<y\),如果 \(a\) 对 \(a[i]、a[j]\) 有影响,即:\(a[i]、a[j]\) 本来不是逆序对,但对 \(a\) 取异或后成为逆序对(或者反过来,本来是逆序对,对 \(a\) 取异或后不是逆序对)那就说明 \(a[i]、a[j]\) 的更高位全部相同,因此 \(b\) 对\(a[i]、a[j]\) 没有影响。
换句话说,不存在 \(a[i]、a[j]\) 既受 \(a\) 影响又受 \(b\) 影响,所以位之间互不影响。
点击查看代码
#include <bits/stdc++.h>
using namespace std;
#define int long long
constexpr int inf = 1e18;
struct Fenwick {
int n;
vector<int> f;
Fenwick(int n_) {
init(n_);
}
void init(int n_) {
n = n_;
f.assign(n, 0);
}
void add(int x, int k) {
for (; x < n; x += x & -x) {
f[x] += k;
}
}
int pre(int x) {
int res = 0;
for (; x; x -= x & -x) {
res += f[x];
}
return res;
}
};
void solve() {
int n;
cin >> n;
vector<int> a(n + 1), pos(n + 1);
Fenwick f(n + 1);
int ans1 = 0;
for (int i = 1; i <= n; i++) {
cin >> a[i];
pos[a[i]] = i;
ans1 += i - 1 - f.pre(a[i] + 1);
f.add(a[i] + 1, 1);
}
int x = 0;
for (int i = 0; (1LL << i) < n; i++) {
int cur = (1LL << i), res = 0;
vector<int> b(a);
f.init(n + 1);
for (int j = 1; j <= n; j++) {
b[j] ^= cur;
res += j - 1 - f.pre(b[j] + 1);
f.add(b[j] + 1, 1);
}
if (res < ans1) {
x ^= cur;
}
}
int ans2 = 1;
f.init(n + 1);
for (int i = 1; i <= n; i++) {
a[i] ^= x;
ans2 += i - 1 - f.pre(a[i] + 1);
f.add(a[i] + 1, 1);
}
cout << min(ans1, ans2) << "\n";
}
signed main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int _ = 1;
// cin >> _;
while (_--) {
solve();
}
return 0;
}
Problem B - Eliminate Tree
2023 ICPC Southeastern Europe Regional Contest E
对一棵树进行以下两种操作:
- 操作一:加点
- 操作二:取边 \((u,v)\) ,保证 \(deg(u)=1\),\(deg(v) \leq 2\),删除点 \(u\) 和点 \(v\)
要删除所有点,求最小操作次数。
考虑树上 dp,定义 \(dp[u][0]\) 表示删除 \(u\) 子树中其他点但不删除 \(u\) 的最小操作次数,\(dp[u][1]\) 表示删除 \(u\) 子树中所有点的最小操作次数。
则有(\(v\) 为 \(u\) 子节点):
点击查看代码
#include<bits/stdc++.h>
using namespace std;
const int N = 2e5+10;
vector<int> e[N];
int n, dp[N][2];
void dfs(int u, int fa){
if(u != 1 && e[u].size() == 1){
dp[u][0] = 0;
dp[u][1] = 2;
return;
}
dp[u][0] = 0; dp[u][1] = 2e9;
for(auto v: e[u]) if(v != fa){
dfs(v, u);
dp[u][0] += dp[v][1];
}
for(auto v: e[u]) if(v != fa){
dp[u][1] = min(dp[u][1], dp[u][0]-dp[v][1]+dp[v][0]+1);
}
}
int main(){
ios::sync_with_stdio(0);
cin >> n;
if(n == 1){
cout << "2\n";
return 0;
}
for(int i = 1; i < n; i++){
int u, v; cin >> u >> v;
e[u].push_back(v);
e[v].push_back(u);
}
dfs(1, 0);
cout << dp[1][1] << endl;
return 0;
}
Problem C - Yet Another Segments Subset
CF 1399F Yet Another Segments Subset
给出 \(n\) 个整数区间,问最多能从中取出多少个区间,使得这些区间彼此包含或互不相交(\(n \sim 3000\))。
考虑区间 DP,记状态 \(dp_{l,r}\) 表示区间 \([l,r]\) 中最多能选出多少个满足要求的区间。对于 \(dp_{l,r}\) ,首先需要判断是否有范围恰为 \([l,r]\) 的原始区间,如果有,则直接计入答案。之后我们针对区间的边界考虑转移,对于左边界 \(l\)(当然也可以选择右边界),分以下两种情况讨论:
- 与左边界有关,状态 \([l,r]\) 由状态 \([l,m]\) 和 \([m+1,r]\) 合并而来
- 与左边界无关,直接继承上一个状态 \([l+1,r]\)
因此有如下转移方程:
其中 \(e\in\{0,1\}\),表示是否有恰好为 \([l,r]\) 的原始区间。
直接枚举 \(m\) 的话为 \(O(n^3)\) 的复杂度,需要进一步优化。注意到对于区间 \([l,r]\),若原始区间中不存在右边界为 \(r+1\) 的区间,则有 \(dp_{l,r}=dp_{l,r+1}\) ,左边界同理。因此转移的时候只需考虑原始区间中存在的右边界即可,用 \(rg_l\) 记录所有左边界为 \(l\) 的区间的右边界,则 \(m\) 只需考虑 \(rg_l\) 中的值即可。
由于区间的范围较大,DP 前还需要离散化将区间的取值范围压缩到 \(2n\)。
虽然此优化明显是有效的,但直觉上这个做法仍是 \(O(n^3)\) 的,下面我们来分析其时间复杂度。对于每一个(压缩过后的)原始区间 \([ls_i,rs_i]\) ,只有当左边界 \(l=ls_i\) 时才会被枚举到,因此该区间共被枚举 \(2n-rs_i\) 次,为 \(O(n)\) 级别,又因为共有 \(n\) 个原始区间,因此总的时间复杂度为 \(O(n^2)\) 。
Bonus:更多区间计数相关问题
- https://atcoder.jp/contests/abc410/tasks/abc410_g 最大区间套
- https://codeforces.com/contest/2133/problem/F 最小区间覆盖
点击查看代码
#include<bits/stdc++.h>
using namespace std;
using ll = long long;
const int N = 6005;
bool vis[N][N], e[N][N];
int n, m, dp[N][N], l[N], r[N], t[N];
vector<int> v, rg[N];
void clear(){
for(int i = 1; i <= m; i++){
rg[i].clear();
for(int j = 1; j <= m; j++)
dp[i][j] = vis[i][j] = e[i][j] = 0;
}
v.clear();
}
void discrete(){
m = 0;
sort(v.begin(), v.end());
int last = -1;
for(auto x: v){
if(x != last) t[++m] = x;
last = x;
}
for(int i = 1; i <= n; i++){
l[i] = lower_bound(t+1, t+m+1, l[i])-t;
r[i] = lower_bound(t+1, t+m+1, r[i])-t;
e[l[i]][r[i]] = 1;
rg[l[i]].push_back(r[i]);
}
for(int i = 1; i <= m; i++)
sort(rg[i].begin(), rg[i].end());
}
void dfs(int l, int r){
if(vis[l][r] || l > r) return;
vis[l][r] = 1;
dp[l][r] = e[l][r];
if(l < r){
dfs(l+1, r);
dp[l][r] = max(dp[l][r], dp[l+1][r]+e[l][r]);
}
for(auto rt: rg[l]){
if(rt >= r) break;
dfs(l, rt);
dfs(rt+1, r);
dp[l][r] = max(dp[l][r], dp[l][rt]+dp[rt+1][r]+e[l][r]);
}
}
int main(){
ios::sync_with_stdio(0);
cin.tie(0);
int T; cin >> T;
while(T--){
cin >> n;
for(int i = 1; i <= n; i++){
cin >> l[i] >> r[i];
v.push_back(l[i]);
v.push_back(r[i]);
}
discrete();
dfs(1, m);
cout << dp[1][m] << endl;
clear();
}
return 0;
}
Problem D - \(K\) Subsequences
2023 ICPC Southeastern Europe Regional Contest K
令 \(f(a) = t\),则至少存在一个子序列 \(a_i\),使得 \(f(a_i) \ge \lceil \frac{t}{k} \rceil\),因此答案至少是 \(\lceil \frac{t}{k} \rceil\)。下面证明这总是可以做到的。
对每个子序列 \(a_{i}\),维护其后缀最大子段和 \(suf_{i}\)。若在子序列 \(a_{i}\) 末尾加入 \(1\),则 \(suf_{i} \rightarrow suf_i + 1\),加入 \(-1\) 则 \(suf_{i} \rightarrow\max(0, suf_{i} - 1)\)。
依次加入 \(a\) 中的每个元素,若当前元素是 \(1\) 则加入到 \(suf_{i}\) 最小的子序列末尾,是 \(-1\) 则加入到 \(suf_{i}\) 最大的子序列末尾。这样所有的 \(suf_i\) 至多相差 \(1\),并且始终有 \(\sum_{i = 1}^{k} suf_i \le cur\),其中 \(cur\) 是已加入的 \(a\) 的元素的最大子段和,因此我们不会使 \(suf_i > \lceil \frac{t}{k} \rceil\)。
时间复杂度 \(O(n\log k)\)。
点击查看代码
#include <bits/stdc++.h>
using namespace std;
inline int read() {
int x = 0, f = 1;
char c = getchar();
while (!isdigit(c)) {
if (c == '-') {
f = -1;
}
c = getchar();
}
while (isdigit(c)) {
x = (x << 3) + (x << 1) + (c ^ 48);
c = getchar();
}
return f * x;
}
void solve() {
int n = read(), k = read();
int l = 0, r = 1;
vector<int> a(n + 1);
for (int i = 1; i <= n; i++) {
// cin >> a[i];
a[i] = read();
r += (a[i] == 1);
}
vector<int> ans(n + 1, 1);
set<pair<int,int> >s;
vector<int> pre(n + 1);
vector<int> premn(n + 1);
auto check = [&]() {
s.clear();
pre.assign(n + 1, 0);
premn.assign(n + 1, 0);
for(int i=1;i<=k;i++)
{
s.emplace(make_pair(0,i));
}
int cnt1 = 0, res = 0;
int num=0;
for (int i = 1; i <= n; i++) {
if (a[i] == -1) {
auto u = *s.rbegin();
s.erase(prev(s.end()));
ans[i]=u.second;
pre[u.second]=pre[u.second]-1;
premn[u.second]=min(premn[u.second],pre[u.second]);
s.emplace(make_pair(pre[u.second]-premn[u.second],u.second));
}
else
{
auto u = *s.begin();
s.erase(s.begin());
ans[i]=u.second;
pre[u.second]=pre[u.second]+1;
premn[u.second]=min(premn[u.second],pre[u.second]);
s.emplace(make_pair(pre[u.second]-premn[u.second],u.second));
}
}
return true;
};
if (r != 1) check();
for (int i = 1; i <= n; i++) {
printf("%d ", ans[i]);
}
printf("\n");
}
signed main() {
// ios::sync_with_stdio(false);
// cin.tie(nullptr);
// int _ = 1;
// cin >> _;
int _ = read();
while (_--) {
solve();
}
return 0;
}
Problem E - Graph Race
2023 ICPC Southeastern Europe Regional Contest G
给定有 \(n\) 个顶点和 \(m\) 条边的无权无向连通图,每个点 \(u\) 有参数 \(a_u, b_u\) 。对每个与点 \(1\) 有直接连边的点 \(v\) ,求 \(\max_{u \not = v} \{a_u - b_u \cdot dist(u, v) \}\) ,其中 \(dist(u, v)\) 是 \(u\) 到 \(v\) 的最短路。
首先,\(dist(u, v)\) 只会是 \(3\) 种情况之一:
- \(dist(u, v) = dist(1, u) - 1\)
- \(dist(u, v) = dist(1, u)\)
- \(dist(u, v) = dist(1, u) + 1\)
令 \(f(u, x) = a_u - b_u \cdot (dist(1, u) + x)\) 。由于 \(b_u > 0\) ,有 \(f(u, -1) \ge f(u, 0) \ge f(u, 1)\) 。因此如果可以使用 \(dist(1, u) - 1\) 条边到达 \(u\) ,\(f(u, -1), f(u, 0), f(u, 1)\) 一定都可以用来更新答案。
对于任意与点 \(1\) 相邻的点 \(v\) ,一定能够使用 \(dist(1, u) + 1\) 条边到达点 \(u\) 。因此初始化 \(ans_v = \max_{u \not = v} f(u, 1)\) 。
要从点 \(v\) 使用 \(dist(1, u) - 1\) 条边到达点 \(u\) ,我们只能使用 \(dist(1, y) = dist(1, x) + 1\) 的“有向边” \(x \rightarrow y\) ,称这种边为前向边。
要从点 \(v\) 使用 \(dist(1, u)\) 条边到达点 \(u\) ,我们只能使用恰好一条 \(dist(1, y) = dist(1, x)\) 的边 \((x, y)\) ,称这种边为横叉边,其他边必须为前向边。
因此,我们可以建一个有 \(2n\) 个顶点的有向图,建两种边:
- 对原图的前向边 \(x \rightarrow y\) ,建边 \(x \rightarrow y, x + n \rightarrow y + n\)
- 对原图的横叉边 \((x, y)\) ,建边 \(x \rightarrow y + n, y \rightarrow x + n\)
如果能在新图中从 \(v\) 到达 \(u (u \le n)\) ,那么 \(dist(v, u) = dist(1, u) - 1\) ,可以用 \(f(u, -1)\) 更新答案。
如果能在新图中从 \(v\) 到达 \(u (u > n)\) ,那么可以用 \(f(u - n, 0)\) 更新答案。
由于新图是有向无环图,可以用 \(dp\) 得到所有答案。需要注意的是,对于与节点 \(1\) 直接相连的点 \(v\),需要特判答案为 \(a_v-b_v\cdot 0\) 或 \(a_v-b_v\cdot 1\) 的情况。
总的时间复杂度为 \(O(n + m)\)。
点击查看代码
#include<bits/stdc++.h>
using namespace std;
using ll = long long;
using pii = pair<int, int>;
using pil = pair<int, ll>;
using pli = pair<ll, int>;
const int N = 6e5+10;
vector<int> e[N], e1[N];
int n, m, d[N];
ll a[N], b[N], dp[N], dp0[N], f[N];
bool vis[N];
pii edge[N];
pli low[N];
void bfs(int st){
queue<int> q;
q.push(1);
vis[1] = 1;
d[1] = 0;
while(q.size()){
int x = q.front();
q.pop();
for(auto y: e[x]) if(!vis[y]){
vis[y] = 1;
d[y] = d[x]+1;
q.push(y);
}
}
}
void dfs(int x){
if(vis[x]) return;
vis[x] = 1;
if(x > n) dp[x] = max(dp[x], a[x-n]-d[x-n]*b[x-n]);
else dp[x] = max(dp[x], a[x]-(d[x]-1)*b[x]);
for(auto y: e1[x]){
dfs(y);
dp0[x] = max(dp0[x], dp[y]);
}
dp[x] = max(dp[x], dp0[x]);
}
int main(){
ios::sync_with_stdio(0);
cin >> n >> m;
for(int i = 1; i <= n; i++)
cin >> a[i] >> b[i];
for(int i = 1; i <= m; i++){
int u, v; cin >> u >> v;
if(u > v) swap(u, v);
e[u].push_back(v);
e[v].push_back(u);
edge[i] = {u, v};
}
bfs(1);
for(int i = 1; i <= m; i++){
auto &[u, v] = edge[i];
if(d[u] > d[v]) swap(u, v);
if(d[v] == d[u]+1)
e1[u].push_back(v), e1[u+n].push_back(v+n);
else if(d[v] == d[u])
e1[u].push_back(v+n), e1[v].push_back(u+n);
}
for(int i = 1; i <= n; i++)
low[i] = {a[i]-(d[i]+1)*b[i], i};
sort(low+1, low+n+1, greater());
fill(dp+1, dp+2*n+1, -1e18);
fill(dp0+1, dp0+2*n+1, -1e18);
fill(vis+1, vis+2*n+1, 0);
dfs(1);
for(auto v: e[1])
if(dp[v] == a[v]-(d[v]-1)*b[v]) dp[v] = dp0[v];
vector<pil> ans;
ll mx1 = low[1].first, id1 = low[1].second, mx2 = low[2].first;
for(auto v: e[1]){
ll val = dp[v];
if(id1 == v) val = max(val, mx2);
else val = max(val, mx1);
ans.emplace_back(v, val);
}
sort(ans.begin(), ans.end());
for(auto [id, val]: ans) cout << val << "\n";
return 0;
}
Problem F - Christmas Sky
2023 ICPC Southeastern Europe Regional Contest C
首先考虑二分最大距离 \(d\) ,则对于新照片上的任意点 \(\mathbf{p}\) 和旧照片上的任意点 \(\mathbf{q}\) ,都有
其中 \(\mathbf{t}\) 是新照片的平移向量。
改写一下式子,有:
这说明,如果把 \(\mathbf{t}\) 视作点,\(-\mathbf{p}\) 视作 \(\mathbf{q}\) 的平移向量,那么平移后的点到 \(\mathbf{t}\) 的距离一定小于等于 \(d\) 。
换言之,以 \(\mathbf{t}\) 为圆心,\(d\) 为半径的圆一定覆盖了所有的点 \(\mathbf{q}-\mathbf{p}\) 。因此求 \(d\) 的最小值,只要求所有的 \(\mathbf{q}-\mathbf{p}\) 的最小圆覆盖即可。时间复杂度为 \(O(nm)\)。
模板:最小圆覆盖
点击查看代码
#include <bits/stdc++.h>
using namespace std;
#define int long long
constexpr int inf = 1e18;
constexpr int N = 1e5 + 9;
constexpr int mod = 1e9 + 7;
using i128 = __int128_t;
using ld = double;
constexpr ld PI = 3.1415926535;
constexpr ld EPS = 1e-6;
template<class T>
struct Point {
T x, y;
Point(const T &x_ = 0, const T &y_ = 0) : x(x_), y(y_) {}
template<class U>
operator Point<U>() {
return Point<U>(U(x), U(y));
}
friend Point operator+(const Point &a, const Point &b) {
return Point(a.x + b.x, a.y + b.y);
}
friend Point operator-(const Point &a, const Point &b) {
return Point(a.x - b.x, a.y - b.y);
}
friend Point<ld> operator*(const Point &a, ld t) {
return Point<ld>(a.x * t, a.y * t);
}
friend Point<ld> operator/(const Point &a, ld t) {
return Point<ld>(a.x / t, a.y / t);
}
friend bool operator==(const Point &a, const Point &b) {
return a.x == b.x && a.y == b.y;
}
friend istream &operator>>(istream &is, Point &a) {
return is >> a.x >> a.y;
}
friend ostream &operator<<(ostream &os, const Point &a) {
return os << "(" << a.x << ", " << a.y << ")";
}
};
template<class T>
T dot(const Point<T> &a, const Point<T> &b) {
return a.x * b.x + a.y * b.y;
}
template<class T>
T cross(const Point<T> &a, const Point<T> &b) {
return a.x * b.y - a.y * b.x;
}
template<class T>
T square(const Point<T> &a) {
return dot(a, a);
}
template<class T>
ld length(const Point<T> &a) {
return sqrtl(square(a));
}
template<class T>
ld distance(const Point<T> &a, const Point<T> &b) {
return length(a - b);
}
template<class T>
struct Circle {
Point<T> c;
T r;
Circle(Point<T> c_, T r_ = -1) : c(c_), r(r_) {}
};
using P = Point<int>;
mt19937_64 rng(chrono::steady_clock::now().time_since_epoch().count());
Circle<ld> CircleFrom3(P a, P b, P c) {
ld D = 2 * (a.x * (b.y - c.y) + b.x * (c.y - a.y) + c.x * (a.y - b.y));
if (abs(D) < EPS) {
ld d01 = distance(a, b), d02 = distance(a, c), d12 = distance(b, c);
if (d02 <= d01 && d12 <= d01) {
return Circle(Point<ld>(a + b) / 2.0, d01 / 2.0);
} else if (d01 <= d02 && d12 <= d02) {
return Circle(Point<ld>(a + c) / 2.0, d02 / 2.0);
} else {
return Circle(Point<ld>(b + c) / 2.0, d12 / 2.0);
}
}
ld da = dot(a, a), db = dot(b, b), dc = dot(c, c);
Point<ld> center(0, 0);
center.x = 1.0 * (da * (b.y - c.y) + db * (c.y - a.y) + dc * (a.y - b.y)) / D;
center.y = 1.0 * (da * (c.x - b.x) + db * (a.x - c.x) + dc * (b.x - a.x)) / D;
ld radius = distance<ld>(center, a);
return Circle(center, radius);
}
Circle<ld> MinimumEnclosingCircle(vector<P> &p) {
int n = p.size();
if (n == 0) {
return {{0, 0}, -1};
} else if (n == 1) {
return {p[0], 0};
}
Circle<ld> mec = {Point<ld>(p[0] + p[1]) / 2.0, distance<ld>(p[0], p[1]) / 2.0};
for (int i = 2; i < n; i++) {
if (square(Point<ld>(p[i]) - mec.c) < mec.r * mec.r + EPS) {
continue;
}
mec = {Point<ld>(p[0] + p[i]) / 2.0, distance<ld>(p[0], p[i]) / 2.0};
for (int j = 1; j < i; j++) {
if (square(Point<ld>(p[j]) - mec.c) < mec.r * mec.r + EPS) {
continue;
}
mec = {Point<ld>(p[i] + p[j]) / 2.0, distance<ld>(p[i], p[j]) / 2.0};
for (int k = 0; k < j; k++) {
if (square(Point<ld>(p[k]) - mec.c) < mec.r * mec.r + EPS) {
continue;
}
mec = CircleFrom3(p[i], p[j], p[k]);
}
}
}
return mec;
}
void solve() {
int n;
cin >> n;
vector<array<int, 2>> a(n);
for (int i = 0; i < n; i++) {
cin >> a[i][0] >> a[i][1];
}
vector<P> p;
int m;
cin >> m;
for (int i = 0; i < m; i++) {
int x, y;
cin >> x >> y;
for (auto [xx, yy] : a) {
p.push_back({x - xx, y - yy});
}
}
shuffle(p.begin(), p.end(), rng);
auto cir = MinimumEnclosingCircle(p);
cout << fixed << setprecision(10) << cir.r << " " << cir.c.x << " " << cir.c.y << "\n";
}
signed main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int _ = 1;
// cin >> _;
while (_--) {
solve();
}
return 0;
}
Problem G - AND-OR closure
2023 ICPC Southeastern Europe Regional Contest A
给出 \(n\) 个正整数,求它们进行 \(\&\) 和 \(|\) 运算得到的闭包的大小,即两两进行与或位运算一共能得到多少个不同的数。
综合性非常强的一道题,考察了图论(序理论)、位运算、计数以及卡常等多方面知识。符号说明:\(x^{[i]}\) 表示在二进制下 \(x\) 的第 \(i\) 位。
首先我们可以考虑去除一些没有用的位,不难发现以下两种位存在冗余:
- 若所有数的第 \(i\) 位都相同,则该位是无效的
- 对于第 \(i\) 位和第 \(j\) 位,所有数的这两位都相同,只需保留一位即可
处理完冗余的位之后,我们考虑各位之间的约束关系。对于第 \(i\) 位和第 \(j\) 位,如果所有第 \(i\) 位为 \(1\) 的数其第 \(j\) 位都为 \(1\),我们则称有偏序关系 \(i\longrightarrow j\)。这等价于将所有位视为节点建图,在满足上述关系的位 \(i\) 和 \(j\) 之间添加有向边,于是我们得到一张关于位的 DAG。当我们通过这张图构造数时,若第 \(i\) 位填 \(1\),由于上述约束,从 \(i\) 出发所能到达的所有位都要填 \(1\) 。记所有位构成集合 \(\mathbb{S}\) ,我们称上述集合为 \(\mathbb{S}\) 的一个闭合子集;节点 \(i\) 没有入度,称为该闭合子集的一个极小元。要求原集合闭包的大小,也就等价于统计闭合子集的数量。
我们先简单证明一下上述构造方法是合法的。记位 \(i\) 对应的最小闭合子集为 \(S_i\)(即选了 \(i\) 之后必选的那些位),这些位构成正整数 \(x_i\)。记 \(P_i\) 表示满足 \(a_k^{[i]}=1 (k\in [1,n])\) 的数构成的集合,则将集合 \(P_i\) 中所有的数按位与就能得到 \(x_i\)。这是因为对于 \(j\in S_i\) ,任意 \(y\in P_i\) 都满足 \(y^{[j]}=1\);对于 \(j\notin S_i\),则必有 \(y\in P_i\) 使得 \(y^{[j]}=0\),否则将存在偏序关系 \(i\longrightarrow j\),这与 \(j\notin S_i\) 矛盾。通过按位或,可以求得不同闭合子集之间的并,也就构造出了 \(\mathbb{S}\) 的所有闭合子集。
直接统计一张图的闭合子集可能并不好做,但我们可以考虑求其反链。反链是指相互之间都不能到达的点集,其与闭合子集一一对应:
- 对于闭合子集 \(S\),其所有的极小元构成一个反链,若极小元 \(i\) 与 \(j\) 之间存在 \(i\) 到 \(j\) 的路径,则子图外还有 \(i\longrightarrow k\) 未被考虑,这与闭合子集的定义矛盾
- 对于反链 \(T\),从 \(T\) 中的节点出发取走所有能到达的点,即可唯一地构造出一个闭合子集
因此原问题转换为求一张 DAG 的反链数量。记共有 \(m\) 个有效位,最朴素的做法是枚举所有 \(k\in [0,2^m)\),从 \(k\) 中任取两位 \(i\) 和 \(j\) (保证 \(k^{[i]}=k^{[j]}=1\)),检查是否存在路径 \(i\longrightarrow j\) 或 \(j\longrightarrow i\),若存在,则状态 \(k\) 不是反链。该方法的复杂度为 \(O(m\cdot 2^m)\),检查路径是否存在可以通过 Floyd 预处理。但由于 \(m\sim 40\),直接暴力算是不行的,我们可以通过类似于双向搜索的方法来进行优化。
我们先按上述方法处理前 \(m/2\) 位,对于后 \(m/2\) 位考虑与前一半的匹配。枚举 \(k_0\in [0,2^{m/2})\),令\(k=2^{m/2}+k_0\),首先判断状态 \(k\) 是否为反链,如果不是则直接跳过。记 \(k\) 对于前半部分的最大匹配为 \(mask\) ,初始为全 \(1\);枚举 \(k\) 的每一位 \(i\in [m/2,m)\),对于 \(k^{[i]}=1\),若 \(i\) 与 \(j\in [0, m/2)\) 之间存在路径,则有 \(mask^{[j]}=0\) 。对于前 \(m/2\) 位,\(mask\) 的所有子集都是可取的。因此我们只需对前半部分进行 SOS DP ,\(dp[mask]\) 即为我们需要的值。
最终的复杂度为 \(O(n\log^2(\max a_i)+\sqrt{\max a_i} \log(\max a_i))\) ,最多要跑大约 \(3\times 10^8\) 次,但原题时限只给了 \(1\sec\),需要注意卡常,可以采用 bitset 等方法进行优化。
下面这份代码跑了约 \(800\space\mathrm{ms}\),并不是很理想,如果有更好的写法欢迎交流。
点击查看代码
#include<bits/stdc++.h>
using namespace std;
using ll = long long;
const int N = 2e5+5, M = 45, INF = 0x3f3f3f3f;
int n, m = 40, e[M][M], d[M][M];
ll a[N], b[N], dp[1<<20];
bitset<50> ban;
bitset<N> bs[50], _bs;
struct DSU{
int n;
vector<int> fa;
DSU(int _n = 0): n(_n){
fa.resize(n);
iota(fa.begin(), fa.end(), 0);
}
int find(int x){
if(fa[x] == x) return x;
return fa[x] = find(fa[x]);
}
void merge(int x, int y){
x = find(x), y = find(y);
if(x > y) swap(x, y);
fa[y] = x;
}
}dsu(40);
inline int reset_bit(int x, int n, int m){
return x&(((1<<n)-1)^(1<<m));
}
int main(){
ios::sync_with_stdio(0);
cin.tie(0);
cin >> n;
for(int i = 1; i <= n; i++)
cin >> a[i];
for(int i = 0; i < 40; i++){
int cnt = 0;
for(int j = 1; j <= n; j++)
cnt += a[j]>>i&1;
if(cnt == n || cnt == 0){
ban[i] = 1;
continue;
}
for(int j = i+1; j < 40; j++){
bool eq = 1;
for(int k = 1; k <= n; k++)
if((a[k]>>i&1) != (a[k]>>j&1)) {eq = 0; break;}
if(eq) dsu.merge(i, j);
}
}
for(int i = 0; i < 40; i++){
if(i != dsu.find(i)) ban[i] = 1;
m -= ban[i];
}
for(int i = 1; i <= n; i++){
b[i] = 0;
for(int j = 0, k = 0; j < 40; j++)
if(!ban[j]) b[i] |= (a[i]>>j&1)<<k, k++;
}
for(int i = 0; i < m; i++)
for(int k = 1; k <= n; k++)
bs[i][k] = b[k]>>i&1;
for(int i = 0; i < m; i++){
for(int j = i+1; j < m; j++){
_bs = bs[i]&bs[j];
if(_bs == bs[i]) e[i][j] = 1;
if(_bs == bs[j]) e[j][i] = 1;
}
}
memset(d, 0x3f, sizeof(d));
for(int k = 0; k < m; k++)
for(int i = 0; i < m; i++)
for(int j = 0; j < m; j++){
if(i == j) {d[i][j] = 0; continue;}
if(e[i][j]) d[i][j] = 1;
d[i][j] = min(d[i][j], d[i][k]+d[k][j]);
}
int hf = m>>1;
for(int k = 0; k < (1<<hf); k++){
dp[k] = 1;
for(int i = 0; k>>i; i++)
for(int j = i+1; k>>j; j++){
if((k>>i&1)+(k>>j&1) < 2) continue;
if(d[i][j] < INF || d[j][i] < INF) dp[k] = 0;
}
}
for(int j = 0; j < hf; j++)
for(int i = 0; i < (1<<hf); i++)
if(i>>j&1) dp[i] += dp[i^(1<<j)];
ll ans = 0;
for(int k = 0; k < (1<<m-hf); k++){
bool valid = 1;
for(int i = 0; k>>i; i++)
for(int j = i+1; k>>j; j++){
if((k>>i&1)+(k>>j&1) < 2) continue;
if(d[i+hf][j+hf] < INF || d[j+hf][i+hf] < INF) valid = 0;
}
if(!valid) continue;
int mask = (1ll<<hf)-1;
for(int i = 0; k>>i; i++){
if(!(k>>i&1)) continue;
for(int j = 0; j < hf; j++)
if(d[i+hf][j] < INF || d[j][i+hf] < INF) mask = reset_bit(mask, hf, j);
}
ans += dp[mask];
}
cout << ans << endl;
return 0;
}
浙公网安备 33010602011771号