Fair and Square
Codeforces Round 1107 (Div. 3) E
문제
링크: Fair and Square
난이도: 1600
풀이
편의상 을 트리의 루트로 정하겠습니다.
트리에서 세 정점을 잡고 개의 경로를 그려보면 세 경로가 모두 지나치는 한 점을 발견할 수 있습니다.
나아가 임의의 세 정점 , , 를 골랐을 때, 트리의 루트를 적당히 조절해서 가 되도록 할 수 있습니다.
라 합시다.
가 의 서브트리 내부에 있는 경우 가 의 서브트리 내부에 있다면 를 루트로, 의 내부에 있다면 로 잡으면 됩니다.
가 의 서브트리 바깥에 있다면 을 루트로 삼으면 됩니다.
이렇게 조절한 루트를 이라 합시다.
그렇다면 로 바꿔 쓸 수 있고
이 됩니다.
이 수가 제곱수가 되려면 이 제곱수여야 합니다.
따라서 이 문제는 이 제곱수인 에 대해 , , 를 고르는 경우의 수를 구하는 문제가 됩니다.
이 다르면 다른 triplet이므로 루트 을 리루팅을 이용해 바꿔가면서 경우를 모두 계산해주면 답을 구할 수 있습니다.
경우의 수 계산
의 자식 개를 골라서 각 서브트리 내에서 정점을 하나씩 고르거나, 을 고른 후 의 자식 개의 서브트리 내에서 정점을 하나씩 고르면 good triplet이 됩니다.
의 자식의 수를 라 할 때 이 경우의 수를 Naive하게 계산하면 이지만 수학을 이용하면 에 계산할 수 있습니다.
정점 의 서브트리의 크기를 라 합시다. 의 자식 개를 골라서 각 서브트리에서 정점을 하나씩 고르는 경우는 입니다. 이므로 에 계산할 수 있습니다. 이 값을 라고 합시다.
그럼 를 이용해서 자식 개를 고르는 경우도 로 계산할 수 있습니다.
여기에 을 고르고 나머지 두개를 자식 서브트리에서 고르는 개도 포함해서 더해주면 됩니다.
리루팅
위에서 보았다시피 루트 에서 경우의 수를 계산하기 위해서는 각 자식마다 서브트리 크기를 알고 있어야 합니다.
의 자식 로 루트를 옮기는 경우 의 서브트리 크기를 로 갱신해주면 됩니다. 는 루트가 되므로 으로 갱신하면 됩니다.
총 시간 복잡도는 입니다.
코드
#include <bits/stdc++.h>
using namespace std;
using ll = long long;
bool isPerfect(int x) {
int r = round(sqrt(x));
return r * r == x;
}
void solve() {
int n;
cin >> n;
vector<int> a(n + 1);
for (int i = 1; i <= n; i++) cin >> a[i];
vector<vector<int>> adj(n + 1);
for (int i = 0; i < n - 1; i++) {
int u, v;
cin >> u >> v;
adj[u].push_back(v);
adj[v].push_back(u);
}
vector<ll> sz(n + 1);
auto calcSz = [&](auto self, int v, int p) -> void {
sz[v] = 1;
for (int c : adj[v]) {
if (c == p) continue;
self(self, c, v);
sz[v] += sz[c];
}
};
calcSz(calcSz, 1, 0);
ll ans = 0;
auto rerooting = [&](auto self, int v, int p) -> void {
for (int c : adj[v]) {
if (c == p) continue;
sz[v] = n - sz[c];
sz[c] = n;
self(self, c, v);
sz[c] = n - sz[v];
}
if (not isPerfect(a[v])) return;
auto comb2 = [&](int c) { return sz[c] * (n - 1 - sz[c]); };
ll B = 0;
for (int c : adj[v]) B += comb2(c);
B /= 2;
ll T = 0;
for (int c : adj[v]) T += sz[c] * (B - comb2(c));
T /= 3;
ans += B + T;
};
rerooting(rerooting, 1, 0);
cout << ans << '\n';
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int t;
cin >> t;
while (t--) solve();
return 0;
}댓글
이름과 이메일을 입력해 댓글을 남겨주세요.