스크류바가 코딩하는 블로그
카테고리아카이브태그소개
자동

© 2026 스쿠루. All rights reserved.

개인정보 처리방침RSS

Fair and Square

2026.08.13|1 min read|PS

Codeforces Round 1107 (Div. 3) E

Codeforces수학조합론정수론그래프 이론그래프 탐색트리DP트리 DP전방향 트리 DP

문제

링크: Fair and Square

난이도: 1600

풀이

편의상 111을 트리의 루트로 정하겠습니다.

트리에서 세 정점을 잡고 333개의 경로를 그려보면 세 경로가 모두 지나치는 한 점을 발견할 수 있습니다.

나아가 임의의 세 정점 uuu, vvv, www를 골랐을 때, 트리의 루트를 적당히 조절해서 LCA⁡(u,v)=LCA⁡(v,w)=LCA⁡(w,u)\operatorname{LCA}(u, v) = \operatorname{LCA}(v, w) = \operatorname{LCA}(w, u)LCA(u,v)=LCA(v,w)=LCA(w,u)가 되도록 할 수 있습니다.

L=LCA⁡(u,v)L = \operatorname{LCA}(u, v)L=LCA(u,v)라 합시다.

www가 LLL의 서브트리 내부에 있는 경우 www가 uuu의 서브트리 내부에 있다면 uuu를 루트로, vvv의 내부에 있다면 vvv로 잡으면 됩니다.

www가 LLL의 서브트리 바깥에 있다면 LLL을 루트로 삼으면 됩니다.

이렇게 조절한 루트를 rrr이라 합시다.

그렇다면 p(u,v)=p(u,r)⋅p(v,r)rp(u, v) = \dfrac{p(u, r) \cdot p(v, r)}{r}p(u,v)=rp(u,r)⋅p(v,r)​로 바꿔 쓸 수 있고

p(u,v)⋅p(v,w)⋅p(w,u)=[p(u,r)]2⋅[p(v,r)]2⋅[p(w,r)]2r3p(u, v) \cdot p(v, w) \cdot p(w, u) = \frac{[p(u, r)]^2 \cdot [p(v, r)]^2 \cdot [p(w, r)]^2}{r^3}p(u,v)⋅p(v,w)⋅p(w,u)=r3[p(u,r)]2⋅[p(v,r)]2⋅[p(w,r)]2​

이 됩니다.

이 수가 제곱수가 되려면 rrr이 제곱수여야 합니다.

따라서 이 문제는 ara_rar​이 제곱수인 rrr에 대해 uuu, vvv, www를 고르는 경우의 수를 구하는 문제가 됩니다.

rrr이 다르면 다른 triplet이므로 루트 rrr을 리루팅을 이용해 바꿔가면서 경우를 모두 계산해주면 답을 구할 수 있습니다.

경우의 수 계산

rrr의 자식 333개를 골라서 각 서브트리 내에서 정점을 하나씩 고르거나, rrr을 고른 후 rrr의 자식 222개의 서브트리 내에서 정점을 하나씩 고르면 good triplet이 됩니다.

rrr의 자식의 수를 CCC라 할 때 이 경우의 수를 Naive하게 계산하면 O(C3)O(C^3)O(C3)이지만 수학을 이용하면 O(C)O(C)O(C)에 계산할 수 있습니다.

정점 vvv의 서브트리의 크기를 svs_vsv​라 합시다. rrr의 자식 222개를 골라서 각 서브트리에서 정점을 하나씩 고르는 경우는 ∑si(∑sj−si)2\dfrac{\sum s_i(\sum s_j - s_i)}{2}2∑si​(∑sj​−si​)​입니다. ∑sj=n−1\sum s_j = n - 1∑sj​=n−1이므로 O(C)O(C)O(C)에 계산할 수 있습니다. 이 값을 BBB라고 합시다.

그럼 BBB를 이용해서 자식 333개를 고르는 경우도 ∑si(B−si(∑sj−si))3\dfrac{\sum s_i(B - s_i(\sum s_j - s_i))}{3}3∑si​(B−si​(∑sj​−si​))​로 계산할 수 있습니다.

여기에 rrr을 고르고 나머지 두개를 자식 서브트리에서 고르는 BBB개도 포함해서 더해주면 됩니다.

리루팅

위에서 보았다시피 루트 rrr에서 경우의 수를 계산하기 위해서는 각 자식마다 서브트리 크기를 알고 있어야 합니다.

rrr의 자식 ccc로 루트를 옮기는 경우 rrr의 서브트리 크기를 n−scn - s_cn−sc​로 갱신해주면 됩니다. ccc는 루트가 되므로 nnn으로 갱신하면 됩니다.


총 시간 복잡도는 O(n+∑C)=O(n)O(n + \sum C) = O(n)O(n+∑C)=O(n)입니다.

코드

#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;
}

목차

  • 문제
  • 풀이
  • 코드

댓글

이름과 이메일을 입력해 댓글을 남겨주세요.