rt,本人用的线段树二分,结果 70 分,求调
#include <bits/stdc++.h>
#define ll long long
using namespace std;
const int N = 1e6 + 1;
int n, k, a[N], rd[N], rt;
ll ans;
vector<int> g[N];
bool vis[N];
struct SegmentTree
{
int Or[N << 2];
int ls(int x) { return x << 1; }
int rs(int x) { return (x << 1) | 1; }
void init() { memset(Or, 0, sizeof Or); }
SegmentTree() { init(); }
void modify(int x, int l, int r, int pos, int val)
{
if (l == r)
return Or[x] = val, void();
int mid = (l + r) >> 1;
if (pos <= mid)
modify(ls(x), l, mid, pos, val);
else
modify(rs(x), mid + 1, r, pos, val);
Or[x] = Or[ls(x)] | Or[rs(x)];
}
int query(int x, int l, int r, int k, int pre)
{
if (Or[1] < k)
return 0;
if (l == r)
return l;
int mid = (l + r) >> 1;
if ((pre | Or[rs(x)]) >= k)
return query(rs(x), mid + 1, r, k, pre);
return query(ls(x), l, mid, k, pre | Or[rs(x)]);
}
} sgt;
template <typename T>
void read(T &x)
{
x = 0;
T f = 1;
char c = getchar();
for (; !isdigit(c); c = getchar())
if (c == '-')
f = -1;
for (; isdigit(c); c = getchar())
x = (x << 3) + (x << 1) + c - 48;
x *= f;
}
template <typename T>
void write(T x)
{
if (x > 9)
write(x / 10);
putchar(x % 10 + 48);
}
template <typename T>
void print(T x, char ed = '\n')
{
if (x < 0)
putchar('-'), x = -x;
write(x), putchar(ed);
}
void dfs(int u, int dep)
{
sgt.modify(1, 1, n, dep, a[u]), ans += 1ll * sgt.query(1, 1, n, k, 0);
for (int v : g[u])
dfs(v, dep + 1);
sgt.modify(1, 1, n, dep, 0);
}
signed main()
{
read(n), read(k);
for (int i = 1; i <= n; i++)
read(a[i]);
for (int i = 1, u, v; i < n; i++)
read(u), read(v), g[u].emplace_back(v), vis[v] = 1;
for (int i = 1; i <= n; i++)
{
if (!vis[i])
{
rt = i;
break;
}
}
dfs(rt, 1), print(ans);
return 0;
}